diff --git a/.github/actions/action.yml b/.github/actions/action.yml index 069ac4395df..97f6c17c8fd 100644 --- a/.github/actions/action.yml +++ b/.github/actions/action.yml @@ -77,7 +77,7 @@ runs: run: echo "node_name=$NODE_NAME" | tee -a "$GITHUB_OUTPUT" - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ inputs.sha }} @@ -94,7 +94,7 @@ runs: sudo chown -R $(whoami) /home/runner/ 2>/dev/null || true - name: Setup python - uses: actions/setup-python@v5 + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5 with: python-version: '3.12' @@ -296,7 +296,7 @@ runs: fi - name: Upload coverage - uses: actions/upload-artifact@v6 + uses: actions/upload-artifact@b7c566a772e6b6bfb58ed0dc250532a479d7789f # v6 if: ${{ always() && steps.check.outputs.coverage_report != 'none' }} with: name: ${{ steps.check.outputs.coverage_report }} @@ -307,7 +307,7 @@ runs: - name: Upload logs id: upload-logs - uses: actions/upload-artifact@v6 + uses: actions/upload-artifact@b7c566a772e6b6bfb58ed0dc250532a479d7789f # v6 if: always() continue-on-error: true with: @@ -321,7 +321,7 @@ runs: run: sleep 10 - name: Retry log upload - uses: actions/upload-artifact@v6 + uses: actions/upload-artifact@b7c566a772e6b6bfb58ed0dc250532a479d7789f # v6 if: ${{ always() && steps.upload-logs.outcome == 'failure' }} with: name: ${{ steps.check.outputs.logs_report }}-retry diff --git a/.github/copy-pr-bot.yaml b/.github/copy-pr-bot.yaml index 10e02754f21..9d8510f443c 100644 --- a/.github/copy-pr-bot.yaml +++ b/.github/copy-pr-bot.yaml @@ -1,4 +1,4 @@ enabled: true auto_sync_draft: false auto_sync_ready: true -trustees_override: ["AAnoosheh", "ArEsKay3", "Autumn1998", "BestJuly", "BoxiangW", "CarlosGomes98", "ChenhanYu", "Connor-XY", "FDecaYed", "HaochenYuan", "HollowMan6", "ISEEKYAN", "JRD971000", "Leili", "Mellonta", "Phlip79", "QiZhangNV", "RPrenger", "ShriyaRishab", "Victarry", "WanZzzzzz", "Wohox", "YangFei1990", "ZhiyuLi-Nvidia", "adistomar", "ahmadki", "aklife97", "alokpathy", "ananthsub", "anlthms", "aroshanghias-nvd", "ashehper", "asolergi-nv", "athitten", "balasaajay", "buptzyb", "chtruong814", "cjld", "cspades", "cuichenx", "deepakn94", "desh2608", "dimapihtar", "dingqingy-nv", "duncanriach", "erhoo82", "ericharper", "fanshiqing", "faradawn", "fitsumreda", "freewym", "frsun-nvda", "gautham-kollu", "gdengk", "goelarushi", "guihong-nv", "guyueh1", "hexinw-nvidia", "huvunvidia", "hxbai", "ilml", "jalbericiola", "janEbert", "jaredcasper", "jenchen13", "jiaji-huang", "jiemingz", "jingqiny-99", "jkamalu", "jon-barker", "jstjohn", "kajalj22", "kamran-nvidia", "kevalmorabia97", "kingformatty", "ko3n1g", "ksivaman", "kunlunl", "kvareddy", "kwyss-nvidia", "lauradang", "layalir", "lhb8125", "liding-nv", "lmcafee-nvidia", "maanug-nv", "macandro96", "mathemakitten", "matthieule", "mchrzanowski", "mehraakash", "minitu", "mkhona-nvidia", "nanz-nv", "ntajbakhsh", "parthmannan", "philipcmonk", "prajwal1210", "pthombre", "rapatel", "rhewett-nv", "rogerwaleffe", "sajadn", "sanandaraj5597", "sancha", "santhnm2", "sbak5", "shanmugamr1992", "sharathts", "sheliang-nv", "shengf-nv", "shifangx", "shjwudp", "sidsingh-nvidia", "skyw", "sraman-rgb", "sudhakarsingh27", "svcnemo-autobot", "tdene", "theothermike", "thomasdhc", "tomlifu", "trintamaki", "tylerpoon", "wdykas", "wplf", "wujingyue", "xiaoyao0115", "xuantengh", "xuwchen", "yaox12", "yaoyu-33", "yashaswikarnati", "yeyu-nvidia", "yobibyte", "youngeunkwon0405", "yqwangustc", "yueshen2016", "yuzhongw-nvidia", "zhehuaichen", "zhongbozhu"] +trustees_override: ["AAnoosheh", "ArEsKay3", "Autumn1998", "BestJuly", "BoxiangW", "CarlosGomes98", "ChenhanYu", "Connor-XY", "FDecaYed", "HaochenYuan", "HollowMan6", "ISEEKYAN", "JF-D", "JRD971000", "Leili", "Mellonta", "Phlip79", "QiZhangNV", "RPrenger", "ShriyaRishab", "Victarry", "WanZzzzzz", "Wohox", "YangFei1990", "ZhiyuLi-Nvidia", "adistomar", "ahmadki", "aklife97", "alokpathy", "ananthsub", "anlthms", "aroshanghias-nvd", "ashehper", "asolergi-nv", "athitten", "balasaajay", "buptzyb", "chtruong814", "cjld", "cspades", "cuichenx", "deepakn94", "desh2608", "dimapihtar", "dingqingy-nv", "duncanriach", "ehosseiniasl", "erhoo82", "ericharper", "fanshiqing", "faradawn", "fitsumreda", "freewym", "frsun-nvda", "gautham-kollu", "gdengk", "goelarushi", "guihong-nv", "guyueh1", "hexinw-nvidia", "huvunvidia", "hxbai", "ilml", "jalbericiola", "janEbert", "jaredcasper", "jenchen13", "jiaji-huang", "jiemingz", "jingqiny-99", "jkamalu", "jon-barker", "jstjohn", "kajalj22", "kamran-nvidia", "kevalmorabia97", "kevjshih", "kingformatty", "ko3n1g", "ksivaman", "kunlunl", "kvareddy", "kwyss-nvidia", "lauradang", "layalir", "lhb8125", "liding-nv", "lmcafee-nvidia", "maanug-nv", "macandro96", "mathemakitten", "matthieule", "mchrzanowski", "mehraakash", "minitu", "mkhona-nvidia", "nanz-nv", "ntajbakhsh", "nvcsathe", "parthmannan", "philipcmonk", "prajwal1210", "pthombre", "rapatel", "rhewett-nv", "rogerwaleffe", "sajadn", "sanandaraj5597", "sancha", "santhnm2", "sbak5", "shanmugamr1992", "sharathts", "sheliang-nv", "shengf-nv", "shifangx", "shjwudp", "sidsingh-nvidia", "skyw", "sraman-rgb", "sudhakarsingh27", "svcnemo-autobot", "tdene", "theothermike", "thomasdhc", "tomlifu", "trintamaki", "tylerpoon", "wdykas", "wplf", "wujingyue", "xiaoyao0115", "xuantengh", "xuwchen", "yaox12", "yaoyu-33", "yashaswikarnati", "yeyu-nvidia", "yobibyte", "youngeunkwon0405", "yqwangustc", "yueshen2016", "yuzhongw-nvidia", "zhehuaichen", "zhongbozhu"] diff --git a/.github/oncall_schedule.json b/.github/oncall_schedule.json index 7ba2c00c095..ba5e92c5980 100644 --- a/.github/oncall_schedule.json +++ b/.github/oncall_schedule.json @@ -1,16 +1,4 @@ [ - { - "user": "cspades", - "date": "2026-07-08" - }, - { - "user": "dimapihtar", - "date": "2026-07-15" - }, - { - "user": "guihong-nv", - "date": "2026-07-22" - }, { "user": "ilml", "date": "2026-07-29" @@ -46,5 +34,17 @@ { "user": "cspades", "date": "2026-09-23" + }, + { + "user": "dimapihtar", + "date": "2026-09-30" + }, + { + "user": "guihong-nv", + "date": "2026-10-07" + }, + { + "user": "ilml", + "date": "2026-10-14" } ] diff --git a/.github/workflows/_build_test_publish_wheel.yml b/.github/workflows/_build_test_publish_wheel.yml index 9e37e068b6d..b6849318a55 100644 --- a/.github/workflows/_build_test_publish_wheel.yml +++ b/.github/workflows/_build_test_publish_wheel.yml @@ -43,7 +43,7 @@ jobs: PUBLISH_DRYRUN: ${{ inputs.dry-run }} steps: - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ inputs.ref }} @@ -141,7 +141,7 @@ jobs: " - name: Upload wheels - uses: actions/upload-artifact@v6 + uses: actions/upload-artifact@b7c566a772e6b6bfb58ed0dc250532a479d7789f # v6 with: name: wheels-${{ matrix.PACKAGE }}-${{ matrix.PLATFORM }}-${{ inputs.dry-run && 'dry-run' || 'release' }} path: dist/ @@ -165,7 +165,7 @@ jobs: PACKAGE: ${{ matrix.PACKAGE }} steps: - name: Download wheels - uses: actions/download-artifact@v7 + uses: actions/download-artifact@37930b1c2abaa49bbe596cd826c3c89aef350131 # v7 with: name: wheels-${{ matrix.PACKAGE }}-${{ matrix.PLATFORM }}-${{ inputs.dry-run && 'dry-run' || 'release' }} path: dist/ diff --git a/.github/workflows/_update_dependencies.yml b/.github/workflows/_update_dependencies.yml index b8410f8fc00..d899855a2df 100644 --- a/.github/workflows/_update_dependencies.yml +++ b/.github/workflows/_update_dependencies.yml @@ -33,7 +33,7 @@ jobs: TARGET_BRANCH: ${{ inputs.target-branch }} steps: - name: Checkout repo - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ env.TARGET_BRANCH }} @@ -60,7 +60,7 @@ jobs: fi - name: Checkout repo - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ env.SOURCE_BRANCH }} @@ -77,7 +77,7 @@ jobs: bash -c 'uv lock --upgrade' - name: Upload lock file - uses: actions/upload-artifact@v6 + uses: actions/upload-artifact@b7c566a772e6b6bfb58ed0dc250532a479d7789f # v6 with: name: lock-file-${{ env.SOURCE_BRANCH }} path: uv.lock @@ -90,7 +90,7 @@ jobs: TARGET_BRANCH: ${{ inputs.target-branch }} steps: - name: Checkout code - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: token: ${{ secrets.PAT }} ref: ${{ env.TARGET_BRANCH }} @@ -103,12 +103,12 @@ jobs: fi - name: Download lock file - uses: actions/download-artifact@v7 + uses: actions/download-artifact@37930b1c2abaa49bbe596cd826c3c89aef350131 # v7 with: name: lock-file-${{ env.SOURCE_BRANCH }} - name: Create Bump PR - uses: peter-evans/create-pull-request@v8 + uses: peter-evans/create-pull-request@5f6978faf089d4d20b00c7766989d076bb2fc7f1 # v8 id: create-pull-request env: title: "chore(beep boop 🤖): Bump `uv.lock` (${{ inputs.target-branch}}) (${{ needs.pre-flight.outputs.date }})" diff --git a/.github/workflows/auto-assign-milestone.yml b/.github/workflows/auto-assign-milestone.yml index b972329bac1..f3ee6709a29 100644 --- a/.github/workflows/auto-assign-milestone.yml +++ b/.github/workflows/auto-assign-milestone.yml @@ -18,7 +18,7 @@ jobs: - name: Get PR info id: get-pr-info if: startsWith(github.ref, 'refs/heads/pull-request/') - uses: nv-gha-runners/get-pr-info@main + uses: nv-gha-runners/get-pr-info@090577647b8ddc4e06e809e264f7881650ecdccf # main - name: Check if PR has milestone id: check_milestone diff --git a/.github/workflows/auto-reminder-bot.yml b/.github/workflows/auto-reminder-bot.yml index 72a48e9539e..23460c4e7dc 100644 --- a/.github/workflows/auto-reminder-bot.yml +++ b/.github/workflows/auto-reminder-bot.yml @@ -14,10 +14,10 @@ jobs: if: github.repository == 'NVIDIA/Megatron-LM' steps: - name: Check out repository code - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Set up Python - uses: actions/setup-python@v6 + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 with: python-version: "3.10" diff --git a/.github/workflows/auto-swap-labels.yml b/.github/workflows/auto-swap-labels.yml index d38fb65d210..9bc7c701fc7 100644 --- a/.github/workflows/auto-swap-labels.yml +++ b/.github/workflows/auto-swap-labels.yml @@ -32,7 +32,7 @@ jobs: id: get-pr if: github.event_name == 'workflow_run' continue-on-error: true - uses: actions/download-artifact@v7 + uses: actions/download-artifact@37930b1c2abaa49bbe596cd826c3c89aef350131 # v7 with: name: pr-number path: pr-number @@ -54,11 +54,11 @@ jobs: - name: Check out repository code if: steps.pr.outputs.number - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Set up Python if: steps.pr.outputs.number - uses: actions/setup-python@v6 + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 with: python-version: "3.10" diff --git a/.github/workflows/auto-update-copy-pr-bot.yml b/.github/workflows/auto-update-copy-pr-bot.yml index 07fdcfbfbb8..d05a844dc0b 100644 --- a/.github/workflows/auto-update-copy-pr-bot.yml +++ b/.github/workflows/auto-update-copy-pr-bot.yml @@ -11,7 +11,7 @@ jobs: if: github.repository == 'NVIDIA/Megatron-LM' steps: - name: Checkout code - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: token: ${{ secrets.PAT }} ref: main diff --git a/.github/workflows/cherry-pick-release-commit.yml b/.github/workflows/cherry-pick-release-commit.yml index 9da305f07e6..2dcef2a06cd 100644 --- a/.github/workflows/cherry-pick-release-commit.yml +++ b/.github/workflows/cherry-pick-release-commit.yml @@ -20,7 +20,7 @@ on: jobs: cherry-pick: - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_cherry_pick.yml@v0.65.9 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_cherry_pick.yml@0cb71cd98aa47ba338d8e38514387d5ceecfedff # v0.65.9 if: github.repository == 'NVIDIA/Megatron-LM' with: target-branches-pattern: 'core_(*dev_)?r[0-9]+\.[0-9]+\.[0-9]+' diff --git a/.github/workflows/cicd-approve-test-queue.yml b/.github/workflows/cicd-approve-test-queue.yml index 32b82a66e19..120b40f4fbb 100644 --- a/.github/workflows/cicd-approve-test-queue.yml +++ b/.github/workflows/cicd-approve-test-queue.yml @@ -30,10 +30,10 @@ jobs: contributor_type: [internal, external] steps: - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Set up Python - uses: actions/setup-python@v6 + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 with: python-version: "3.12" diff --git a/.github/workflows/cicd-main.yml b/.github/workflows/cicd-main.yml index ff39b026c1b..0f0b8b900ea 100644 --- a/.github/workflows/cicd-main.yml +++ b/.github/workflows/cicd-main.yml @@ -54,14 +54,14 @@ jobs: DISABLE_EXTERNAL_CONTRIBUTOR: ${{ vars.DISABLE_EXTERNAL_CONTRIBUTOR }} steps: - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: token: ${{ env.GITHUB_TOKEN }} - name: Get PR info id: get-pr-info if: startsWith(github.ref, 'refs/heads/pull-request/') && github.event_name == 'push' - uses: nv-gha-runners/get-pr-info@main + uses: nv-gha-runners/get-pr-info@090577647b8ddc4e06e809e264f7881650ecdccf # main - name: Check NVIDIA SSO membership id: check-sso @@ -133,7 +133,7 @@ jobs: pre-flight: needs: [is-not-external-contributor] if: github.repository == 'NVIDIA/Megatron-LM' - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_cicd_preflight.yml@v1.0.0 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_cicd_preflight.yml@6a2f81195fd910ae91d3c001d2e32ceb6d82e975 # v1.0.0 configure: runs-on: ubuntu-latest @@ -154,7 +154,7 @@ jobs: - name: Get PR info id: get-pr-info if: startsWith(github.ref, 'refs/heads/pull-request/') && github.event_name == 'push' - uses: nv-gha-runners/get-pr-info@main + uses: nv-gha-runners/get-pr-info@090577647b8ddc4e06e809e264f7881650ecdccf # main # Resolve a single SHA used by the build, every test job, and every # downstream checkout so that the container image, golden values, and @@ -341,12 +341,12 @@ jobs: ) steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: fetch-depth: 0 - name: Install uv - uses: astral-sh/setup-uv@v8.1.0 + uses: astral-sh/setup-uv@08807647e7069bb48b6ef5acd8ec9567f424441b # v8.1.0 with: version: 0.7.2 @@ -357,7 +357,27 @@ jobs: - name: Get PR info id: get-pr-info if: startsWith(github.ref, 'refs/heads/pull-request/') && github.event_name == 'push' - uses: nv-gha-runners/get-pr-info@main + uses: nv-gha-runners/get-pr-info@090577647b8ddc4e06e809e264f7881650ecdccf # main + + - name: Validate updated golden values + if: github.event_name == 'merge_group' || (startsWith(github.ref, 'refs/heads/pull-request/') && github.event_name == 'push') + env: + BASE_REF: ${{ github.event.merge_group.base_ref || fromJSON(steps.get-pr-info.outputs.pr-info || '{}').base.ref }} + run: | + BASE_REF="${BASE_REF#refs/heads/}" + git fetch origin "$BASE_REF" + mapfile -t GOLDEN_VALUES_FILES < <( + git diff --name-only --diff-filter=ACMR \ + --merge-base "origin/$BASE_REF" -- \ + ':(glob)tests/functional_tests/test_cases/**/golden_values*.json' + ) + + if (( ${#GOLDEN_VALUES_FILES[@]} == 0 )); then + echo "No golden value files were updated; skipping validation." + exit 0 + fi + + python3 tools/check_golden_values.py "${GOLDEN_VALUES_FILES[@]}" - name: Run linting if: startsWith(github.ref, 'refs/heads/pull-request/') && github.event_name == 'push' @@ -406,7 +426,7 @@ jobs: mbridge-test-suite: ${{ needs.configure.outputs.mbridge_suite }} steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: How-To run: bash .github/scripts/readme.sh @@ -418,9 +438,9 @@ jobs: - configure - cicd-wait-in-queue - cicd-parse-downstream-testing - # skip downstream mbridge testing on PR pushes by - # default. They still run for merge_group and nightly (schedule / - # workflow_dispatch) triggers, and PR authors can opt in by adding the + # Skip downstream MBridge testing for docs-only changes and PR pushes by + # default. Non-docs merge_group and nightly (schedule / workflow_dispatch) + # triggers still run it, and PR authors can opt in by adding the # "Run MBridge tests" label — all three cases set # configure.outputs.run_mbridge == 'true'. if: | @@ -430,6 +450,7 @@ jobs: && needs.cicd-parse-downstream-testing.result != 'cancelled' && vars.ENABLE_CICD_MBRIDGE_TESTING == 'true' && needs.configure.outputs.run_mbridge == 'true' + && needs.pre-flight.outputs.docs_only == 'false' && ( success() || needs.pre-flight.outputs.is_ci_workload == 'true' @@ -441,10 +462,10 @@ jobs: - name: Get PR info id: get-pr-info if: startsWith(github.ref, 'refs/heads/pull-request/') && github.event_name == 'push' - uses: nv-gha-runners/get-pr-info@main + uses: nv-gha-runners/get-pr-info@090577647b8ddc4e06e809e264f7881650ecdccf # main - name: Checkout MBridge and create testing branch - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: main repository: NVIDIA-NeMo/Megatron-Bridge @@ -461,7 +482,7 @@ jobs: git push origin ${{ env.MBRIDGE_BRANCH_NAME }} --force - name: Trigger MBridge tests - uses: convictional/trigger-workflow-and-wait@v1.6.5 + uses: convictional/trigger-workflow-and-wait@f69fa9eedd3c62a599220f4d5745230e237904be # v1.6.5 env: MBRIDGE_BRANCH_NAME: mcore-testing-${{ fromJSON(steps.get-pr-info.outputs.pr-info || '{}').number || github.run_id }} with: @@ -497,7 +518,7 @@ jobs: && (needs.cicd-mbridge-testing.result == 'success' || needs.cicd-mbridge-testing.result == 'failure') steps: - name: Send Slack alert - uses: NVIDIA-NeMo/FW-CI-templates/.github/actions/send-slack-alert@main + uses: NVIDIA-NeMo/FW-CI-templates/.github/actions/send-slack-alert@209ac7913b0419a5ccbac47b02d00fbea4939243 # main with: webhook: ${{ secrets.SLACK_WH_MLM_MB_ALERTS }} message: | @@ -556,15 +577,15 @@ jobs: - name: Get PR info id: get-pr-info if: startsWith(github.ref, 'refs/heads/pull-request/') && github.event_name == 'push' - uses: nv-gha-runners/get-pr-info@main + uses: nv-gha-runners/get-pr-info@090577647b8ddc4e06e809e264f7881650ecdccf # main - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ needs.configure.outputs.sha }} - name: Setup python - uses: actions/setup-python@v6 + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 with: python-version: 3.12 @@ -627,10 +648,10 @@ jobs: fi - name: Set up Docker Buildx - uses: docker/setup-buildx-action@v4.0.0 + uses: docker/setup-buildx-action@4d04d5d9486b7bd6fa91e7baf45bbb4f8b9deedd # v4.0.0 - name: Build and push - uses: docker/build-push-action@v7.1.0 + uses: docker/build-push-action@bcafcacb16a39f128d818304e6c9c0c18556b85f # v7.1.0 with: file: ${{ steps.base-image.outputs.dockerfile }} push: true @@ -671,7 +692,7 @@ jobs: && !cancelled() steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ needs.configure.outputs.sha }} - name: Parse unit tests @@ -717,7 +738,7 @@ jobs: PIP_RETRIES: 5 steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ needs.configure.outputs.sha }} - name: main @@ -757,7 +778,7 @@ jobs: && !cancelled() steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ needs.configure.outputs.sha }} - name: Parse unit tests @@ -805,7 +826,7 @@ jobs: PIP_RETRIES: 5 steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ needs.configure.outputs.sha }} - name: main @@ -896,7 +917,7 @@ jobs: integration-tests-h100: ${{ steps.main.outputs.integration-tests-h100 }} steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ needs.configure.outputs.sha }} @@ -961,7 +982,7 @@ jobs: && needs.cicd-parse-integration-tests-h100.result == 'success' steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ needs.configure.outputs.sha }} - name: main @@ -995,7 +1016,7 @@ jobs: integration-tests-gb200: ${{ steps.main.outputs.integration-tests-gb200 }} steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ needs.configure.outputs.sha }} @@ -1062,7 +1083,7 @@ jobs: && vars.ENABLE_GB200_TESTING == 'true' steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: ref: ${{ needs.configure.outputs.sha }} - name: main @@ -1104,7 +1125,7 @@ jobs: permissions: write-all steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Get workflow result id: result @@ -1214,7 +1235,7 @@ jobs: && github.repository == 'NVIDIA/Megatron-LM' steps: - name: Generate fake coverage report - uses: actions/github-script@v8 + uses: actions/github-script@ed597411d8f924073f98dfc5c65a23a2325f34cd # v8 with: github-token: ${{ secrets.PAT }} script: | @@ -1245,13 +1266,13 @@ jobs: - name: Get PR info id: get-pr-info if: startsWith(github.ref, 'refs/heads/pull-request/') && github.event_name == 'push' - uses: nv-gha-runners/get-pr-info@main + uses: nv-gha-runners/get-pr-info@090577647b8ddc4e06e809e264f7881650ecdccf # main - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Download coverage reports of current branch - uses: actions/download-artifact@v7 + uses: actions/download-artifact@37930b1c2abaa49bbe596cd826c3c89aef350131 # v7 with: pattern: coverage-${{ matrix.flag }}-* @@ -1272,7 +1293,7 @@ jobs: ls -al - name: Upload coverage reports to Codecov - uses: codecov/codecov-action@v5 + uses: codecov/codecov-action@0fb7174895f61a3b6b78fc075e0cd60383518dac # v5 with: token: ${{ secrets.CODECOV_TOKEN }} verbose: true @@ -1280,7 +1301,7 @@ jobs: base_sha: ${{ fromJSON(steps.get-pr-info.outputs.pr-info || '{}').base.sha }} - name: Upload artifacts - uses: actions/upload-artifact@v6 + uses: actions/upload-artifact@b7c566a772e6b6bfb58ed0dc250532a479d7789f # v6 with: name: coverage-${{ matrix.flag }}-aggregated path: | @@ -1301,7 +1322,7 @@ jobs: echo "pr_number=$PR_NUMBER" >> $GITHUB_OUTPUT - name: Comment on PR with action run URL - uses: actions/github-script@v8 + uses: actions/github-script@ed597411d8f924073f98dfc5c65a23a2325f34cd # v8 with: github-token: ${{ secrets.PAT }} script: | diff --git a/.github/workflows/claude-complexity-label.yml b/.github/workflows/claude-complexity-label.yml index 541cdb4e539..d9b765fb3dd 100644 --- a/.github/workflows/claude-complexity-label.yml +++ b/.github/workflows/claude-complexity-label.yml @@ -5,38 +5,39 @@ on: types: [ready_for_review] jobs: - label-complexity: - name: Label PR Complexity + analyze_complexity: + name: Analyze PR Complexity runs-on: ubuntu-latest permissions: contents: read - pull-requests: write - issues: write - id-token: write + pull-requests: read + issues: read + outputs: + label_json: ${{ steps.analyze.outputs.structured_output }} env: - GH_TOKEN: ${{ secrets.PAT }} REPO: ${{ github.repository }} PR_NUMBER: ${{ github.event.pull_request.number }} steps: - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: fetch-depth: 0 - name: Run Claude Complexity Analysis - uses: anthropics/claude-code-action@v1 + id: analyze + uses: anthropics/claude-code-action@be7b93b1907a4abad570368f3c74b6fe3807510b # v1 env: ANTHROPIC_BASE_URL: ${{ secrets.NVIDIA_INFERENCE_URL }} CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS: "1" DISABLE_PROMPT_CACHING: "1" with: anthropic_api_key: ${{ secrets.NVIDIA_INFERENCE_KEY }} - github_token: ${{ secrets.PAT }} + github_token: ${{ github.token }} prompt: | REPO: ${{ env.REPO }} PR NUMBER: ${{ env.PR_NUMBER }} - You are a PR complexity analyzer. Your job is to analyze the diff of this PR and apply exactly one complexity label. + You are a PR complexity analyzer. Your job is to analyze the diff of this PR and return exactly one complexity label. STEPS: 1. Get the PR diff by running: gh pr diff $PR_NUMBER --repo $REPO @@ -47,19 +48,45 @@ jobs: 3. Compute "real code line changes" using this formula: real_code_line_changes = (number of real code lines changed) + (number of test lines changed / 10) Count both added and removed lines. Do not count unchanged context lines. Do not count comments or docstrings. - 4. Remove any previously applied complexity or docs-only labels: - gh pr edit $PR_NUMBER --repo $REPO --remove-label "complexity: low,complexity: medium,complexity: high,docs-only" - 5. Apply exactly ONE label using the gh CLI: - - If there are ZERO real code lines and ZERO test lines (only docs-only changes), apply label "docs-only": - gh pr edit $PR_NUMBER --repo $REPO --add-label "docs-only" - - If real_code_line_changes < 100, apply label "complexity: low": - gh pr edit $PR_NUMBER --repo $REPO --add-label "complexity: low" - - If real_code_line_changes >= 100 and < 500, apply label "complexity: medium": - gh pr edit $PR_NUMBER --repo $REPO --add-label "complexity: medium" - - If real_code_line_changes >= 500, apply label "complexity: high": - gh pr edit $PR_NUMBER --repo $REPO --add-label "complexity: high" + 4. Return exactly ONE label: + - If there are ZERO real code lines and ZERO test lines (only docs-only changes), return "docs-only". + - If real_code_line_changes < 100, return "complexity: low". + - If real_code_line_changes >= 100 and < 500, return "complexity: medium". + - If real_code_line_changes >= 500, return "complexity: high". - Do NOT post any comments on the PR. Only apply the label. + Do NOT post comments, edit the PR, modify labels, or run any write operation. claude_args: | - --allowedTools "Bash(gh pr diff:*),Bash(gh pr edit:*),Bash(gh pr view:*)" + --allowedTools "Bash(gh pr diff:*),Bash(gh pr view:*)" --model "${{ vars.CLAUDE_MODEL }}" + --json-schema '{"type":"object","properties":{"label":{"type":"string","enum":["docs-only","complexity: low","complexity: medium","complexity: high"]}},"required":["label"],"additionalProperties":false}' + + apply-complexity-label: + name: Apply PR Complexity Label + runs-on: ubuntu-latest + needs: analyze_complexity + permissions: + pull-requests: write + issues: write + env: + GH_TOKEN: ${{ github.token }} + REPO: ${{ github.repository }} + PR_NUMBER: ${{ github.event.pull_request.number }} + LABEL_JSON: ${{ needs.analyze_complexity.outputs.label_json }} + steps: + - name: Apply validated complexity label + run: | + set -euo pipefail + + label=$(echo "$LABEL_JSON" | jq -r '.label // empty') + case "$label" in + "docs-only"|"complexity: low"|"complexity: medium"|"complexity: high") + ;; + *) + echo "::error::Claude returned invalid complexity label: $label" + exit 1 + ;; + esac + + gh pr edit "$PR_NUMBER" --repo "$REPO" \ + --remove-label "complexity: low,complexity: medium,complexity: high,docs-only" || true + gh pr edit "$PR_NUMBER" --repo "$REPO" --add-label "$label" diff --git a/.github/workflows/claude-copy-to-main.yml b/.github/workflows/claude-copy-to-main.yml index 24659574b77..dc7b56e1529 100644 --- a/.github/workflows/claude-copy-to-main.yml +++ b/.github/workflows/claude-copy-to-main.yml @@ -5,8 +5,8 @@ on: types: [created] jobs: - copy-to-main: - name: Copy PR to Main + authorize_copy: + name: Authorize Copy to Main if: | github.event_name == 'issue_comment' && github.event.issue.pull_request && @@ -14,14 +14,14 @@ jobs: contains(github.event.comment.body, '/claude copy') runs-on: ubuntu-latest permissions: - contents: write - pull-requests: write issues: write - id-token: write + pull-requests: read env: GH_TOKEN: ${{ secrets.PAT }} REPO: ${{ github.repository }} PR_NUMBER: ${{ github.event.issue.number }} + outputs: + base_ref: ${{ steps.pr-info.outputs.base_ref }} steps: - name: Check commenter has write access env: @@ -34,10 +34,12 @@ jobs: fi - name: Check PR is merged and targets non-main + id: pr-info run: | PR_JSON=$(gh pr view $PR_NUMBER --repo $REPO --json baseRefName,mergedAt) PR_BASE=$(echo "$PR_JSON" | jq -r .baseRefName) PR_MERGED=$(echo "$PR_JSON" | jq -r .mergedAt) + echo "base_ref=$PR_BASE" >> "$GITHUB_OUTPUT" if [ "$PR_BASE" = "main" ]; then gh pr comment $PR_NUMBER --repo $REPO --body "❌ This PR already targets \`main\`. The Claude copy command only works on PRs targeting non-main branches." @@ -49,18 +51,36 @@ jobs: exit 1 fi + prepare_copy: + name: Prepare Copy Patch + runs-on: ubuntu-latest + needs: authorize_copy + permissions: + contents: read + pull-requests: read + issues: read + env: + GH_TOKEN: ${{ github.token }} + REPO: ${{ github.repository }} + PR_NUMBER: ${{ github.event.issue.number }} + COPY_BRANCH: copy-pr-${{ github.event.issue.number }}-to-main + steps: - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: fetch-depth: 0 - token: ${{ secrets.PAT }} - name: Fetch PR head ref from fork run: | git fetch origin pull/$PR_NUMBER/head:pr-$PR_NUMBER-head + - name: Configure Git + run: | + git config user.name "svcnvidia-nemo-ci" + git config user.email "svcnvidia-nemo-ci@nvidia.com" + - name: Run Claude Copy to Main - uses: anthropics/claude-code-action@v1 + uses: anthropics/claude-code-action@be7b93b1907a4abad570368f3c74b6fe3807510b # v1 env: ANTHROPIC_BASE_URL: ${{ secrets.NVIDIA_INFERENCE_URL }} CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS: "1" @@ -68,12 +88,14 @@ jobs: with: anthropic_api_key: ${{ secrets.NVIDIA_INFERENCE_KEY }} trigger_phrase: "/claude copy" - github_token: ${{ secrets.PAT }} + github_token: ${{ github.token }} prompt: | REPO: ${{ env.REPO }} PR NUMBER: ${{ env.PR_NUMBER }} + SOURCE BASE REF: ${{ needs.authorize_copy.outputs.base_ref }} + COPY BRANCH: ${{ env.COPY_BRANCH }} - You are a PR copy assistant. Your job is to apply the final changes from a merged PR onto a new branch based on `main` and create a new PR targeting `main`. + You are a PR copy assistant. Your job is to apply the final changes from a merged PR onto a local branch based on `main`. The PR's commits originated from a fork and have been fetched locally as the branch: pr-${PR_NUMBER}-head @@ -81,19 +103,15 @@ jobs: 1. Get the PR details (title, body, and base branch): gh pr view $PR_NUMBER --repo $REPO --json title,body,baseRefName - 2. Configure git for committing (use the svcnvidia-nemo-ci service account since secrets.PAT belongs to it): - git config user.name "svcnvidia-nemo-ci" - git config user.email "svcnvidia-nemo-ci@nvidia.com" - - 3. Create a new branch from `main`: + 2. Create a new local branch from `main`: git checkout main git pull origin main - git checkout -b copy-pr-${PR_NUMBER}-to-main + git checkout -b $COPY_BRANCH - 4. Generate a patch of the PR's final changes and apply it: + 3. Generate a patch of the PR's final changes and apply it: MERGE_BASE=$(git merge-base origin/ pr-${PR_NUMBER}-head) git diff $MERGE_BASE pr-${PR_NUMBER}-head | git apply --3way - (Replace with the actual base branch name from step 1.) + (Replace with SOURCE BASE REF unless step 1 shows a different base branch.) If the apply fails due to merge conflicts: a. Identify conflicted files: git diff --name-only --diff-filter=U @@ -103,25 +121,101 @@ jobs: without overriding what is already on main. d. Stage the resolved files: git add - 5. Commit the changes: + 4. Commit the changes locally: git add -A - git commit -m "Copy PR #${PR_NUMBER} to main" - - 6. Push the new branch: - git push origin copy-pr-${PR_NUMBER}-to-main - - 7. Create a new PR targeting `main`: - gh pr create --repo $REPO \ - --base main \ - --head copy-pr-${PR_NUMBER}-to-main \ - --title "[Copy to main] " \ - --body "🤖 **This PR was auto-generated by Claude** via the Claude copy workflow.\n\nCherry-picked from #${PR_NUMBER}.\n\n---\n\n" - - 8. Comment on the original PR with a link to the newly created PR. + git commit -s -m "Copy PR #${PR_NUMBER} to main" IMPORTANT: + - Do NOT push. + - Do NOT create a pull request. + - Do NOT comment on the original PR. + - Do NOT use gh for any operation except reading PR metadata. - When resolving merge conflicts, favor `main` over the non-main branch. Do not override changes already on main. - - Do NOT force push. claude_args: | - --allowedTools "Bash(git:*),Bash(gh:*),Read,Edit" + --allowedTools "Bash(git:*),Bash(gh pr view:*),Read,Edit" --model "${{ vars.CLAUDE_MODEL }}" + + - name: Export copy patch + run: | + set -euo pipefail + + git status --short + test "$(git rev-parse --abbrev-ref HEAD)" = "$COPY_BRANCH" + test -z "$(git status --porcelain)" + test "$(git rev-list --count origin/main..HEAD)" -gt 0 + + git diff --binary origin/main..HEAD > "$RUNNER_TEMP/copy-pr.patch" + test -s "$RUNNER_TEMP/copy-pr.patch" + + - name: Upload copy patch + uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4 + with: + name: copy-pr-${{ github.event.issue.number }}-patch + path: ${{ runner.temp }}/copy-pr.patch + if-no-files-found: error + retention-days: 1 + + publish_copy: + name: Publish Copy PR + runs-on: ubuntu-latest + needs: [authorize_copy, prepare_copy] + permissions: + contents: write + pull-requests: write + issues: write + env: + GH_TOKEN: ${{ secrets.PAT }} + REPO: ${{ github.repository }} + PR_NUMBER: ${{ github.event.issue.number }} + COPY_BRANCH: copy-pr-${{ github.event.issue.number }}-to-main + steps: + - name: Checkout repository + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 + with: + fetch-depth: 0 + token: ${{ secrets.PAT }} + + - name: Download copy patch + uses: actions/download-artifact@d3f86a106a0bac45b974a628896c90dbdf5c8093 # v4 + with: + name: copy-pr-${{ github.event.issue.number }}-patch + path: ${{ runner.temp }} + + - name: Create branch, commit, and PR + run: | + set -euo pipefail + + git config user.name "svcnvidia-nemo-ci" + git config user.email "svcnvidia-nemo-ci@nvidia.com" + + git fetch origin main + git checkout -b "$COPY_BRANCH" origin/main + git apply --3way "$RUNNER_TEMP/copy-pr.patch" + git diff --check + git add -A + git commit -s -m "Copy PR #${PR_NUMBER} to main" + git push origin "$COPY_BRANCH" + + PR_JSON=$(gh pr view "$PR_NUMBER" --repo "$REPO" --json title,body) + ORIGINAL_TITLE=$(echo "$PR_JSON" | jq -r '.title') + echo "$PR_JSON" | jq -r '.body // ""' > "$RUNNER_TEMP/original-pr-body.md" + { + echo "🤖 **This PR was auto-generated by Claude** via the Claude copy workflow." + echo + echo "Cherry-picked from #${PR_NUMBER}." + echo + echo "---" + echo + cat "$RUNNER_TEMP/original-pr-body.md" + } > "$RUNNER_TEMP/copy-pr-body.md" + + NEW_PR_URL=$(gh pr create \ + --repo "$REPO" \ + --base main \ + --head "$COPY_BRANCH" \ + --title "[Copy to main] $ORIGINAL_TITLE" \ + --body-file "$RUNNER_TEMP/copy-pr-body.md") + + gh pr comment "$PR_NUMBER" \ + --repo "$REPO" \ + --body "✅ Created copy-to-main PR: $NEW_PR_URL" diff --git a/.github/workflows/claude_review.yml b/.github/workflows/claude_review.yml index 98fe4eac964..c29bbc2df6c 100644 --- a/.github/workflows/claude_review.yml +++ b/.github/workflows/claude_review.yml @@ -32,7 +32,7 @@ jobs: echo "sha=$(gh pr view $PR_NUMBER --repo $REPO --json headRefOid -q .headRefOid)" | tee -a $GITHUB_OUTPUT - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: fetch-depth: 1 ref: ${{ steps.get-pr-head-commit.outputs.sha }} @@ -44,7 +44,7 @@ jobs: -f content='eyes' - name: Run Claude Light Review - uses: anthropics/claude-code-action@v1 + uses: anthropics/claude-code-action@be7b93b1907a4abad570368f3c74b6fe3807510b # v1 env: ANTHROPIC_BASE_URL: ${{ secrets.NVIDIA_INFERENCE_URL }} CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS: "1" @@ -52,7 +52,6 @@ jobs: with: anthropic_api_key: ${{ secrets.NVIDIA_INFERENCE_KEY }} trigger_phrase: "/claude review" - show_full_output: true claude_args: | --allowedTools "mcp__github_inline_comment__create_inline_comment,Bash(gh pr comment:*),Bash(gh pr diff:*),Bash(gh pr view:*),Bash(gh pr review:*),Read" --model "${{ vars.CLAUDE_MODEL }}" @@ -132,7 +131,7 @@ jobs: echo "base_ref=$(echo $PR_DATA | jq -r .baseRefName)" >> $GITHUB_OUTPUT - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: fetch-depth: 1 ref: ${{ steps.pr-info.outputs.sha }} @@ -147,7 +146,7 @@ jobs: -f content='eyes' - name: Run Claude Strict Review - uses: anthropics/claude-code-action@v1 + uses: anthropics/claude-code-action@be7b93b1907a4abad570368f3c74b6fe3807510b # v1 env: ANTHROPIC_BASE_URL: ${{ secrets.NVIDIA_INFERENCE_URL }} CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS: "1" @@ -155,7 +154,6 @@ jobs: with: anthropic_api_key: ${{ secrets.NVIDIA_INFERENCE_KEY }} trigger_phrase: "/claude strict-review" - show_full_output: true claude_args: | --allowedTools "mcp__github_inline_comment__create_inline_comment,Bash(gh pr comment:*),Bash(gh pr diff:*),Bash(gh pr view:*),Bash(gh pr review:*),Bash(git diff:*),Bash(git show:*),Bash(git log:*),Read" --model "${{ vars.CLAUDE_MODEL }}" diff --git a/.github/workflows/close-inactive-issue-pr.yml b/.github/workflows/close-inactive-issue-pr.yml index 7dcac837ba9..9f9377b259f 100644 --- a/.github/workflows/close-inactive-issue-pr.yml +++ b/.github/workflows/close-inactive-issue-pr.yml @@ -19,4 +19,4 @@ on: jobs: close-issues: if: github.repository == 'NVIDIA/Megatron-LM' - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_close_inactive_issue_pr.yml@v0.44.0 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_close_inactive_issue_pr.yml@9e07489b8a6bc533c8792099b012c588f4430298 # v0.44.0 diff --git a/.github/workflows/community-bot.yml b/.github/workflows/community-bot.yml index 1a98ece0f85..47a54ec9264 100644 --- a/.github/workflows/community-bot.yml +++ b/.github/workflows/community-bot.yml @@ -21,7 +21,7 @@ on: jobs: community-bot: - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_community_bot.yml@v0.65.10 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_community_bot.yml@f0dadfd1b2d5c3f48a24ded127abd50afbf8ce11 # v0.65.10 with: community_project_id: ${{ vars.COMMUNITY_PROJECT_ID }} if: github.repository == 'NVIDIA/Megatron-LM' diff --git a/.github/workflows/community-request-assignee.yml b/.github/workflows/community-request-assignee.yml index a344690c0f2..a0581ecfd5e 100644 --- a/.github/workflows/community-request-assignee.yml +++ b/.github/workflows/community-request-assignee.yml @@ -126,13 +126,13 @@ jobs: ISSUE_AUTHOR: ${{ github.event.issue.user.login }} steps: - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: fetch-depth: 0 - name: Analyze issue owner with Claude id: claude-analysis - uses: anthropics/claude-code-action@v1 + uses: anthropics/claude-code-action@be7b93b1907a4abad570368f3c74b6fe3807510b # v1 env: ANTHROPIC_BASE_URL: ${{ secrets.NVIDIA_INFERENCE_URL }} CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS: "1" @@ -242,7 +242,7 @@ jobs: - name: Checkout repository if: steps.still-unassigned.outputs.skip != 'true' - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Install assignment dependencies if: steps.still-unassigned.outputs.skip != 'true' diff --git a/.github/workflows/copyright-check.yml b/.github/workflows/copyright-check.yml index 484a66fb0e0..c5a20f9c066 100644 --- a/.github/workflows/copyright-check.yml +++ b/.github/workflows/copyright-check.yml @@ -24,7 +24,7 @@ on: jobs: pre-flight: - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_cicd_preflight.yml@v1.0.0 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_cicd_preflight.yml@6a2f81195fd910ae91d3c001d2e32ceb6d82e975 # v1.0.0 if: github.repository == 'NVIDIA/Megatron-LM' copyright-check: @@ -34,7 +34,7 @@ jobs: || needs.pre-flight.outputs.is_merge_group == 'true' || needs.pre-flight.outputs.is_deployment_workflow == 'true') && github.repository == 'NVIDIA/Megatron-LM' - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_copyright_check.yml@v1.0.0 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_copyright_check.yml@6a2f81195fd910ae91d3c001d2e32ceb6d82e975 # v1.0.0 copyright-check-summary: needs: [pre-flight, copyright-check] @@ -49,7 +49,7 @@ jobs: runs-on: ubuntu-latest steps: - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Result env: diff --git a/.github/workflows/install-test.yml b/.github/workflows/install-test.yml index 3505937cd92..1a1ae490bc9 100644 --- a/.github/workflows/install-test.yml +++ b/.github/workflows/install-test.yml @@ -29,7 +29,7 @@ on: jobs: pre-flight: - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_cicd_preflight.yml@v1.0.0 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_cicd_preflight.yml@6a2f81195fd910ae91d3c001d2e32ceb6d82e975 # v1.0.0 if: github.repository == 'NVIDIA/Megatron-LM' pip-test-pytorch: @@ -49,7 +49,7 @@ jobs: python-version: ["3.12"] steps: - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Set PATH run: | @@ -65,7 +65,7 @@ jobs: run: bash docker/common/install.sh --environment dev --base-image pytorch --python-version ${{ matrix.python-version }} - name: Checkout check-imports - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: repository: NVIDIA-NeMo/FW-CI-templates ref: v0.63.2 @@ -100,7 +100,7 @@ jobs: python-version: ["3.12"] steps: - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Set PATH run: | @@ -119,7 +119,7 @@ jobs: # NGC PyTorch 25.05 has a version of triton that is broken on CPU only machines. # - name: Checkout check-imports - # uses: actions/checkout@v6 + # uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 # with: # repository: NVIDIA-NeMo/FW-CI-templates # ref: v0.63.2 @@ -145,7 +145,7 @@ jobs: && github.repository == 'NVIDIA/Megatron-LM' steps: - name: Checkout - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Get workflow result id: result diff --git a/.github/workflows/nightly-sync-main-to-dev.yml b/.github/workflows/nightly-sync-main-to-dev.yml index 07490d9bade..8b34eb1de0d 100644 --- a/.github/workflows/nightly-sync-main-to-dev.yml +++ b/.github/workflows/nightly-sync-main-to-dev.yml @@ -27,16 +27,13 @@ concurrency: cancel-in-progress: false permissions: - contents: write - pull-requests: write - issues: write - id-token: write + contents: read jobs: # Re-dispatch scheduled runs as workflow_dispatch via a PAT so the heavy # job runs with a real User-type actor. On `schedule` events GitHub sets # `github.actor` to `github-merge-queue` (no Users-API entry), which - # crashes anthropics/claude-code-action@v1 in `checkHumanActor` with a + # crashes anthropics/claude-code-action@be7b93b1907a4abad570368f3c74b6fe3807510b # v1 in `checkHumanActor` with a # 404 before `allowed_bots` is ever consulted. Upstream fix PR # https://github.com/anthropics/claude-code-action/pull/1212 is closed # and unmerged; see issue @@ -64,7 +61,7 @@ jobs: GH_TOKEN: ${{ secrets.PAT }} steps: - name: Checkout repository - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: fetch-depth: 0 token: ${{ secrets.PAT }} @@ -111,14 +108,19 @@ jobs: echo "skip=false" >> "$GITHUB_OUTPUT" fi - - name: Install pre-push merge guard + - name: Install pre-push merge guidance if: steps.check-sync.outputs.skip != 'true' run: | cat > .git/hooks/pre-push <<'HOOK' #!/usr/bin/env bash + + # This hook is advisory. Run the checks in a strict subshell so an + # audit error can be reported without blocking the push. + set +e + ( set -euo pipefail - echo "=== nightly-sync pre-push guard ===" + echo "=== nightly-sync pre-push guidance ===" merge_commit=$(git rev-list --min-parents=2 --max-count=1 HEAD || true) if [ -n "$merge_commit" ]; then @@ -130,8 +132,7 @@ jobs: fi if ! git diff --quiet "$dev_ref" HEAD -- .github/CODEOWNERS; then - echo "ABORT: .github/CODEOWNERS differs from dev. Restore it before pushing." - exit 1 + echo "WARNING: .github/CODEOWNERS differs from dev. Restore it before finalizing the sync." fi for f in pyproject.toml uv.lock docker/Dockerfile.ci.dev; do @@ -148,7 +149,7 @@ jobs: intentional_override_regex='^(megatron/training/training\.py|megatron/training/initialize\.py|megatron/training/utils\.py|megatron/training/datasets/data_samplers\.py|megatron/core/optimizer/layer_wise_optimizer\.py)$' skip_regex='^(pyproject\.toml|uv\.lock|docker/Dockerfile\.ci\.dev|\.github/CODEOWNERS)$' - violations=0 + findings=0 while IFS= read -r f; do [[ "$f" =~ $skip_regex ]] && continue [[ "$f" =~ $intentional_override_regex ]] && continue @@ -164,26 +165,34 @@ jobs: if [ -n "$missing" ]; then echo "=== $f ===" printf '%s\n' "$missing" - violations=$((violations + $(printf '%s\n' "$missing" | grep -c .))) + findings=$((findings + $(printf '%s\n' "$missing" | grep -c .))) fi done < <(git diff --name-only "$dev_ref"..HEAD \ -- '*.py' '*.md' '*.yaml' '*.yml' '*.toml' \ '*.sh' '*.cpp' '*.cu' '*.h' \ | sort -u) - if [ "$violations" -gt 0 ]; then - echo "ABORT: $violations dev-only line(s) were dropped by the merge." - echo "Restore the dev-only code, or document the exact main commit that intentionally removed it." - exit 1 + if [ "$findings" -gt 0 ]; then + echo "WARNING: $findings potential dev-only line removal(s) were detected." + echo "Review each finding: restore merge accidents and document intentional main removals in the PR body." + echo "This audit is advisory; the push will continue." + else + echo "No potential dev-only line removals detected." fi - echo "nightly-sync pre-push guard passed" + echo "nightly-sync pre-push guidance complete" + ) + guidance_status=$? + if [ "$guidance_status" -ne 0 ]; then + echo "WARNING: nightly-sync pre-push guidance failed with status $guidance_status; allowing the push to continue." + fi + exit 0 HOOK chmod +x .git/hooks/pre-push - name: Run Claude Code to merge, fix, and iterate if: steps.check-sync.outputs.skip != 'true' - uses: anthropics/claude-code-action@v1 + uses: anthropics/claude-code-action@be7b93b1907a4abad570368f3c74b6fe3807510b # v1 env: ANTHROPIC_BASE_URL: ${{ secrets.NVIDIA_INFERENCE_URL }} CLAUDE_CODE_DISABLE_EXPERIMENTAL_BETAS: "1" @@ -236,12 +245,14 @@ jobs: conversation in between — that wastes `--max-turns` and creates windows where the agent could forget the loop. - **Pre-push guard:** The workflow installs a local git pre-push - hook that enforces CODEOWNERS, dependency-triple, and dev-feature - preservation checks. You MUST NOT bypass it with `--no-verify`. - If a push fails, read the hook output, restore the dropped dev - code unless main explicitly removed it, and push again only after - the hook passes. + **Pre-push guidance:** The workflow installs a local git pre-push + hook that reports CODEOWNERS, dependency-triple, and dev-feature + preservation findings. It is advisory and MUST NOT block a push. + Do not use `--no-verify`; let the hook run and review its output. + Restore genuine merge accidents and CODEOWNERS changes. For + intentional main removals or formatting/reordering false positives, + document the evidence in the PR body and continue. Do not stop or + ask for authorization solely because advisory findings remain. **Merge strategy:** Start from `origin/dev` and run `git merge origin/main --no-edit`. Do NOT use global @@ -301,7 +312,6 @@ jobs: `Nemo_CICD_Test`, `copyright-check`, `pre-flight`, wheel builds, etc. — is NOT exempt and must reach a terminal green conclusion. - show_full_output: true claude_args: | --allowedTools "Bash,Read,Edit,Write,Grep,Glob,Agent" --model "${{ vars.CLAUDE_MODEL }}" diff --git a/.github/workflows/oncall-assign.yml b/.github/workflows/oncall-assign.yml index 6da0776ffc2..dc96f51b350 100644 --- a/.github/workflows/oncall-assign.yml +++ b/.github/workflows/oncall-assign.yml @@ -30,10 +30,10 @@ jobs: if: ${{ !github.event.pull_request.draft }} steps: - name: Checkout code - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Set up Python - uses: actions/setup-python@v6 + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 with: python-version: '3.10' diff --git a/.github/workflows/oncall-rotation.yml b/.github/workflows/oncall-rotation.yml index 0d5f774e441..66b9fd8ddce 100644 --- a/.github/workflows/oncall-rotation.yml +++ b/.github/workflows/oncall-rotation.yml @@ -28,12 +28,12 @@ jobs: runs-on: ubuntu-latest steps: - name: Checkout code - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: token: ${{ secrets.PAT }} - name: Set up Python - uses: actions/setup-python@v6 + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 with: python-version: "3.10" diff --git a/.github/workflows/release-docs.yml b/.github/workflows/release-docs.yml index 6d619a8a1bc..7207f767522 100644 --- a/.github/workflows/release-docs.yml +++ b/.github/workflows/release-docs.yml @@ -73,7 +73,7 @@ on: jobs: build-docs: - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_build_docs.yml@v0.67.0 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_build_docs.yml@3ab507cd035df3ae37cce8808ed3210ff6e7062b # v0.67.0 with: ref: ${{ inputs.build-docs-ref }} @@ -81,7 +81,7 @@ jobs: runs-on: ubuntu-latest needs: [build-docs] steps: - - uses: actions/checkout@v6 + - uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 with: repository: NVIDIA-NeMo/FW-CI-templates ref: v0.74.0 diff --git a/.github/workflows/release-freeze.yml b/.github/workflows/release-freeze.yml index 8037a8cb4bc..8eccf2caac9 100644 --- a/.github/workflows/release-freeze.yml +++ b/.github/workflows/release-freeze.yml @@ -34,7 +34,7 @@ on: default: true jobs: code-freeze: - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_code_freeze.yml@v1.4.2 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_code_freeze.yml@bfdb5e35067fd8cd91ce21fca4eb1072ffd7ab8c # v1.4.2 with: library-name: Megatron-Core python-package: megatron.core diff --git a/.github/workflows/release.yaml b/.github/workflows/release.yaml index cd193d819eb..1b3d2af292a 100644 --- a/.github/workflows/release.yaml +++ b/.github/workflows/release.yaml @@ -72,7 +72,7 @@ concurrency: jobs: pre-flight: - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_cicd_preflight.yml@v0.94.1 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_cicd_preflight.yml@211c302d648552cecccec610fe796cc22a091f37 # v0.94.1 if: github.repository == 'NVIDIA/Megatron-LM' && github.event_name != 'workflow_dispatch' bump: @@ -83,7 +83,7 @@ jobs: && !(needs.pre-flight.outputs.docs_only == 'true' || needs.pre-flight.outputs.is_merge_group == 'true' || needs.pre-flight.outputs.is_deployment_workflow == 'true') - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_release_bump.yml@v1.4.0 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_release_bump.yml@6dfd1b435cca9e3c2640f7b31c4f37e42c6bf796 # v1.4.0 with: release-branch-pattern: "core_[rv][0-9]*.[0-9]*.[0-9]*" release-ref: ${{ inputs.release-ref || github.sha }} @@ -123,7 +123,7 @@ jobs: github.repository == 'NVIDIA/Megatron-LM' && (success() || !failure()) && !cancelled() - uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_release_finalize.yml@v1.0.0 + uses: NVIDIA-NeMo/FW-CI-templates/.github/workflows/_release_finalize.yml@6a2f81195fd910ae91d3c001d2e32ceb6d82e975 # v1.0.0 with: release-ref: ${{ inputs.release-ref || github.sha }} release-version: ${{ needs.bump.outputs.release-version }} diff --git a/.github/workflows/request-nvskills-ci.yml b/.github/workflows/request-nvskills-ci.yml index 01c9b5c7569..07c0a846c0e 100644 --- a/.github/workflows/request-nvskills-ci.yml +++ b/.github/workflows/request-nvskills-ci.yml @@ -17,6 +17,6 @@ jobs: permissions: contents: read pull-requests: read - uses: NVIDIA/skills/.github/workflows/team-request.yml@main + uses: NVIDIA/skills/.github/workflows/team-request.yml@2528d5b9d3f125c8bc8cf644ea2134adb4322a51 # main secrets: NVSKILLS_CI_DISPATCH_TOKEN: ${{ secrets.NVSKILLS_CI_DISPATCH_TOKEN }} diff --git a/.github/workflows/review-trigger.yml b/.github/workflows/review-trigger.yml index 7375e605aff..e7aabde4113 100644 --- a/.github/workflows/review-trigger.yml +++ b/.github/workflows/review-trigger.yml @@ -22,7 +22,7 @@ jobs: mkdir -p pr echo "${{ github.event.pull_request.number }}" > pr/number - name: Upload PR number - uses: actions/upload-artifact@v6 + uses: actions/upload-artifact@b7c566a772e6b6bfb58ed0dc250532a479d7789f # v6 with: name: pr-number path: pr/ diff --git a/.github/workflows/sync-team-usergroups.yml b/.github/workflows/sync-team-usergroups.yml index 7f32ac55c57..71e1752077e 100644 --- a/.github/workflows/sync-team-usergroups.yml +++ b/.github/workflows/sync-team-usergroups.yml @@ -24,10 +24,10 @@ jobs: runs-on: ubuntu-latest steps: - name: Checkout code - uses: actions/checkout@v6 + uses: actions/checkout@d23441a48e516b6c34aea4fa41551a30e30af803 # v6 - name: Set up Python - uses: actions/setup-python@v6 + uses: actions/setup-python@ece7cb06caefa5fff74198d8649806c4678c61a1 # v6 with: python-version: "3.10" diff --git a/.github/workflows/trigger-mbridge-tests.yml b/.github/workflows/trigger-mbridge-tests.yml index 023851e966a..e828183b322 100644 --- a/.github/workflows/trigger-mbridge-tests.yml +++ b/.github/workflows/trigger-mbridge-tests.yml @@ -25,7 +25,7 @@ jobs: runs-on: ubuntu-latest steps: - name: Trigger MBridge tests - uses: convictional/trigger-workflow-and-wait@v1.6.5 + uses: convictional/trigger-workflow-and-wait@f69fa9eedd3c62a599220f4d5745230e237904be # v1.6.5 with: owner: NVIDIA-NeMo repo: Megatron-Bridge diff --git a/.gitlab-ci.yml b/.gitlab-ci.yml index 2eb1b43be0c..ae27dc6d2f4 100644 --- a/.gitlab-ci.yml +++ b/.gitlab-ci.yml @@ -154,6 +154,7 @@ stages: - integration_tests - functional_tests - publish + - triage default: interruptible: true @@ -268,7 +269,21 @@ variables: - "upgrade-dependencies" description: Type of publish (freeze or final release) + RUN_LINEAR_STATUS: + value: "True" + options: + - "True" + - "False" + description: Reconcile functional-test failures against Linear + RUN_LINEAR_WRITE: + value: "True" + options: + - "True" + - "False" + description: Apply proposed Linear issue opens, updates, and closes + # CI wide variables + NEMO_CI_TRIAGE_CONFIG: .gitlab/nemo-ci-triage.yml CI_MCORE_LTS_IMAGE: ${GITLAB_ENDPOINT}:5005/adlr/megatron-lm/mcore_ci_lts CI_MCORE_DEV_IMAGE: ${GITLAB_ENDPOINT}:5005/adlr/megatron-lm/mcore_ci_dev CI_NEMO_IMAGE: ${GITLAB_ENDPOINT}:5005/adlr/megatron-lm/nemo_ci @@ -282,3 +297,4 @@ include: - .gitlab/stages/03.integration-tests.yml - .gitlab/stages/04.functional-tests.yml - .gitlab/stages/05.publish.yml + - .gitlab/stages/06.triage.yml diff --git a/.gitlab/nemo-ci-triage.yml b/.gitlab/nemo-ci-triage.yml new file mode 100644 index 00000000000..4922d63dac1 --- /dev/null +++ b/.gitlab/nemo-ci-triage.yml @@ -0,0 +1,16 @@ +# Megatron-LM configuration for nemo-ci-triage. + +gitlab: + project_id: 19378 + repo_name: ADLR/megatron-lm + +modules: + megatron_lm: + build_module: megatron-lm + channel_id_env: MCORE_SLACK_CHANNEL_ID + reconcile_proposal: true + team_key: MCORE + project_template: "MCore CI Testing" + enable_linear_open: true + enable_linear_modify: true + enable_linear_close: true diff --git a/.gitlab/scripts/build.sh b/.gitlab/scripts/build.sh index 15c926ed51f..c72bd38a581 100644 --- a/.gitlab/scripts/build.sh +++ b/.gitlab/scripts/build.sh @@ -48,6 +48,11 @@ if [[ -n "$TE_GIT_REF" ]]; then ADDITIONAL_PARAMS+=("--build-arg TE_COMMIT=${TE_GIT_REF}") fi +if [[ "$FILE" == "Dockerfile.linting" ]]; then + ADDITIONAL_PARAMS+=("--build-arg CI_SERVER_URL=${CI_SERVER_URL}") + ADDITIONAL_PARAMS+=("--secret id=NEMO_CI_TRIAGE_TOKEN,env=PAT") +fi + echo $(git rev-parse HEAD) JET_API_VERSION=$(curl -s -u "$ARTIFACTORY_USER:$ARTIFACTORY_TOKEN" "https://sc-hw-artf.nvidia.com/artifactory/api/pypi/hw-joc-pypi/simple/jet-api/" | grep -o 'href="../../jet-api/[0-9.]*/' | sed 's|href="../../jet-api/||;s|/||' | sort -V -r | head -n1) diff --git a/.gitlab/stages/02.test.yml b/.gitlab/stages/02.test.yml index a324ce037fb..93385bd8e1b 100644 --- a/.gitlab/stages/02.test.yml +++ b/.gitlab/stages/02.test.yml @@ -82,6 +82,7 @@ test:unit_tests_configure: "--dependent-job test:unit_tests_configure" "--slurm-account ${CI_SLURM_ACCOUNT}" "--no-enable-warmup" + "--enable-error-extraction" ) - | export PYTHONPATH=$(pwd) @@ -194,8 +195,9 @@ test:unit_tests_notify: fi - export RO_API_TOKEN=${PROJECT_ACCESS_TOKEN_MCORE} - export GITLAB_ENDPOINT - - export TAG_TEAM=$([[ "$CI_COMMIT_BRANCH" == "main" ]] && echo "1" || "0") + - export TAG_TEAM=$([[ "$CI_COMMIT_BRANCH" == "main" ]] && echo "1" || echo "0") - export TEAM_SLUG=$SLACK_ADMIN + - export PYTHONPATH=$(pwd) - | python tests/test_utils/python_scripts/notify.py \ --pipeline-id "${CI_PIPELINE_ID}" \ diff --git a/.gitlab/stages/03.integration-tests.yml b/.gitlab/stages/03.integration-tests.yml index 70fa345e513..603c6e09b52 100644 --- a/.gitlab/stages/03.integration-tests.yml +++ b/.gitlab/stages/03.integration-tests.yml @@ -56,6 +56,7 @@ integration:configure: "--no-enable-warmup" "--dependent-job integration:configure" "--enable-lightweight-mode" + "--enable-error-extraction" ) - | export PYTHONPATH=$(pwd) diff --git a/.gitlab/stages/04.functional-tests.yml b/.gitlab/stages/04.functional-tests.yml index 45b7cc4daab..2671c20282a 100644 --- a/.gitlab/stages/04.functional-tests.yml +++ b/.gitlab/stages/04.functional-tests.yml @@ -87,6 +87,7 @@ functional:configure: "--record-checkpoints ${RECORD_CHECKPOINTS}" "--slurm-account ${CI_SLURM_ACCOUNT}" "--no-enable-warmup" + "--enable-error-extraction" ) - | SMOKE_ARGS=( @@ -101,6 +102,7 @@ functional:configure: "--record-checkpoints false" "--slurm-account ${CI_SLURM_ACCOUNT}" "--no-enable-warmup" + "--enable-error-extraction" ) - | export PYTHONPATH=$(pwd) @@ -426,41 +428,9 @@ functional:run_nemo: allow_failure: true - when: never -functional:smoke_notify: - extends: [.functional_tests_rules] - image: ${UTILITY_IMAGE}:${CI_PIPELINE_ID} - needs: - - functional:smoke-h100 - - functional:smoke-gb200 - tags: - - arch/amd64 - - env/prod - - origin/jet-fleet - - owner/jet-core - - purpose/utility - - team/megatron - script: - - | - if [[ "$CI_COMMIT_BRANCH" == *dev* ]]; then - export WEBHOOK_URL=${MCORE_NOTIFICATION_HOOK_DEV} - else - export WEBHOOK_URL=${MCORE_NOTIFICATION_HOOK} - fi - - export RO_API_TOKEN=${PROJECT_ACCESS_TOKEN_MCORE} - - export GITLAB_ENDPOINT - - | - python tests/test_utils/python_scripts/notify.py \ - --pipeline-id "${CI_PIPELINE_ID}" \ - --check-for smoke-tests \ - --pipeline-context "smoke-${FUNCTIONAL_TEST_SCOPE}" \ - --pipeline-created-at "${CI_PIPELINE_CREATED_AT}" - rules: - - if: $BUILD == "no" - when: never - - if: $FUNCTIONAL_TEST == "yes" && $FUNCTIONAL_TEST_SCOPE =~ /^(mr|nightly)$/ && ($CI_PIPELINE_SOURCE == "schedule" || $CI_COMMIT_BRANCH == "main" || $CI_MERGE_REQUEST_EVENT_TYPE == "merged_result") - when: always - - when: never - +# Sole root Slack notification for functional MR, nightly, weekly, and release +# pipelines. Detailed triage and Linear updates reply in the thread recorded by +# slack_notification.json; individual workload runners must not notify directly. functional:x_notify: extends: [.functional_tests_rules] image: ${UTILITY_IMAGE}:${CI_PIPELINE_ID} @@ -492,19 +462,28 @@ functional:x_notify: - export RO_API_TOKEN=${PROJECT_ACCESS_TOKEN_MCORE} - export GITLAB_ENDPOINT - export CONTEXT=$FUNCTIONAL_TEST_SCOPE - - export TAG_TEAM=$([[ "$CI_COMMIT_BRANCH" == "main" || "$CI_COMMIT_BRANCH" == "dev" ]] && echo "1" || "0") + - export TAG_TEAM=$([[ "$CI_COMMIT_BRANCH" == "main" || "$CI_COMMIT_BRANCH" == "dev" ]] && echo "1" || echo "0") - export TEAM_SLUG=$SLACK_ADMIN + - export PYTHONPATH=$(pwd) - | python tests/test_utils/python_scripts/notify.py \ --pipeline-id "${CI_PIPELINE_ID}" \ --check-for functional-tests \ --pipeline-context $CONTEXT \ - --pipeline-created-at "${CI_PIPELINE_CREATED_AT}" + --pipeline-created-at "${CI_PIPELINE_CREATED_AT}" \ + --summary-output pipeline_summaries.json \ + --failure-buckets-output failure_buckets.json \ + --slack-output slack_notification.json artifacts: when: always paths: - scripts + - pipeline_summaries.json + - failure_buckets.json + - slack_notification.json + - inference_metrics.json + - agent_formatter_debug.txt rules: - if: ($CI_PIPELINE_SOURCE == "schedule" || $CI_COMMIT_BRANCH == "main" || $CI_COMMIT_BRANCH == "dev") && $FUNCTIONAL_TEST == "yes" when: always diff --git a/.gitlab/stages/06.triage.yml b/.gitlab/stages/06.triage.yml new file mode 100644 index 00000000000..b80010540ae --- /dev/null +++ b/.gitlab/stages/06.triage.yml @@ -0,0 +1,110 @@ +.linear_reconcile_rules: + rules: + - if: >- + ($CI_PIPELINE_SOURCE == "schedule" || $CI_COMMIT_BRANCH == "main") && + $FUNCTIONAL_TEST == "yes" && + $RUN_LINEAR_STATUS == "True" + when: always + - when: never + +.linear_triage_job: + stage: triage + image: ${UTILITY_IMAGE}:${CI_PIPELINE_ID} + tags: + - arch/amd64 + - env/prod + - origin/jet-fleet + - owner/jet-core + - purpose/utility + - team/megatron + +triage:linear_reconcile: + extends: [.linear_triage_job, .linear_reconcile_rules] + needs: + - job: functional:x_notify + artifacts: true + script: + - >- + nemo-ci-linear status + --config "${NEMO_CI_TRIAGE_CONFIG}" + --build-module-regex '^megatron-lm$' + --output linear_status_report.json + - >- + nemo-ci-linear reconcile + --failure-buckets failure_buckets.json + --linear-report linear_status_report.json + --pipeline-summaries pipeline_summaries.json + --output linear_action_plan.json + artifacts: + when: always + paths: + - linear_status_report.json + - linear_action_plan.json + - inference_metrics.json + +triage:linear_write: + extends: [.linear_triage_job] + needs: + - job: triage:linear_reconcile + artifacts: true + allow_failure: true + script: + - >- + nemo-ci-linear write + --config "${NEMO_CI_TRIAGE_CONFIG}" + --plan linear_action_plan.json + --output linear_action_plan_post.json + artifacts: + when: always + paths: + - linear_action_plan.json + - linear_action_plan_post.json + - linear_status_report.json + rules: + - if: >- + ($CI_PIPELINE_SOURCE == "schedule" || $CI_COMMIT_BRANCH == "main") && + $FUNCTIONAL_TEST == "yes" && + $RUN_LINEAR_STATUS == "True" && + $RUN_LINEAR_WRITE == "True" + when: always + - when: never + +triage:slack_linear_followup: + extends: [.linear_triage_job] + needs: + - job: functional:x_notify + artifacts: true + - job: triage:linear_write + artifacts: true + allow_failure: true + script: + - >- + nemo-ci-notify + --pipeline-summary slack_notification.json + --linear-plan linear_action_plan_post.json + --slack-bot-token "${MCORE_SLACK_BOT_TOKEN:-${ALERTMANAGER_TOKEN}}" + --slack-channel-id "${MCORE_SLACK_CHANNEL_ID}" + - | + THREAD_TIMESTAMP="$(python -c 'import json; print(json.load(open("slack_notification.json")).get("thread_timestamp") or "")')" + if [[ -z "${THREAD_TIMESTAMP}" ]]; then + echo "No Slack thread timestamp; skipping detailed triage follow-ups." + else + nemo-ci-notify \ + --config "${NEMO_CI_TRIAGE_CONFIG}" \ + --module megatron_lm \ + --only-followup \ + --thread-ts "${THREAD_TIMESTAMP}" \ + --failure-buckets failure_buckets.json \ + --linear-report linear_status_report.json \ + --action-plan linear_action_plan_post.json \ + --slack-bot-token "${MCORE_SLACK_BOT_TOKEN:-${ALERTMANAGER_TOKEN}}" + fi + rules: + # Post the applied Linear actions under the functional-test notification. + - if: >- + ($CI_PIPELINE_SOURCE == "schedule" || $CI_COMMIT_BRANCH == "main") && + $FUNCTIONAL_TEST == "yes" && + $RUN_LINEAR_STATUS == "True" && + $RUN_LINEAR_WRITE == "True" + when: always + - when: never diff --git a/AGENTS.md b/AGENTS.md index e747867f8b0..996c38c17bb 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -24,7 +24,12 @@ skill keyword — infer it from the artifact you read. - All PRs must be created as **drafts**. Use `gh pr create --draft` or the GitHub UI draft option. - Never push branches directly to `https://github.com/NVIDIA/Megatron-LM`. You must push your branch to a personal fork (e.g. `https://github.com//Megatron-LM`), then open a PR from the fork's branch against `NVIDIA/Megatron-LM`. -- Commit PR changes with both `-s` and `-S`: `-s` adds the required `Signed-off-by` trailer, and `-S` signs the commit so copy-pr-bot and `/ok to test` can verify the pushed commit without manually specifying the SHA. +- Commit PR changes with both `-s` and `-S`: `-s` adds the required + `Signed-off-by` trailer, and `-S` signs the commit so copy-pr-bot and `/ok to +test` can verify the pushed commit without manually specifying the SHA. +Megatron Core engineers at NVIDIA should sign using their NVIDIA emails so they +are automatically added to the right user groups on the internal Slack +workspace. - Read @docs/developer/contribute.md for the full contribution policy, including code style, commit message conventions, and issue guidelines. ### Code Quality diff --git a/docker/Dockerfile.ci.lts b/docker/Dockerfile.ci.lts index 6a2042e345d..1e681e2ff81 100644 --- a/docker/Dockerfile.ci.lts +++ b/docker/Dockerfile.ci.lts @@ -46,10 +46,9 @@ COPY megatron/core/package_info.py /workspace/megatron/core/ ENV NVTE_BUILD_NUM_PHILOX_ROUNDS=3 RUN --mount=type=cache,target=/root/.cache/uv \ bash -ex <<"EOF" - export NVTE_CUDA_ARCHS="80;90;100" uv venv ${UV_PROJECT_ENVIRONMENT} --system-site-packages uv sync --only-group build - uv sync --extra mlm --extra ssm --extra te --link-mode copy --locked \ + uv sync --extra mlm --extra ssm --link-mode copy --locked \ --no-install-package torch \ --no-install-package torchvision \ --no-install-package triton \ @@ -72,12 +71,16 @@ EOF # # These used to live in `[project.optional-dependencies].lts` in pyproject.toml, # but were moved out so pyproject.toml can host meaningful per-module -# extras. The pinned set lives in `docker/lts/requirements.txt` and is reviewed -# at LTS bump time only. +# extras. Most of the pinned set lives in `docker/lts/requirements.txt` and is +# reviewed at LTS bump time only. Transformer Engine is installed separately +# because its source extension needs the existing PyTorch/CUDA build environment. COPY docker/lts/requirements.txt /workspace/docker/lts/requirements.txt RUN --mount=type=cache,target=/root/.cache/uv \ bash -ex <<"EOF" + export NVTE_CUDA_ARCHS="80;90;100" uv pip install -r /workspace/docker/lts/requirements.txt + uv pip install --no-build-isolation \ + "transformer-engine @ git+https://github.com/NVIDIA/TransformerEngine.git@b9d690e042b1c4e455214e7dab65d6d3512c05d6" EOF # Install DeepEP diff --git a/docker/Dockerfile.linting b/docker/Dockerfile.linting index bf27b768374..737b8cefac8 100644 --- a/docker/Dockerfile.linting +++ b/docker/Dockerfile.linting @@ -21,3 +21,13 @@ ARG JET_API_VERSION RUN --mount=type=secret,id=JET_INDEX_URLS \ JET_INDEX_URLS=$(cat /run/secrets/JET_INDEX_URLS) && \ uv pip install --no-cache-dir "jet-client~=2.0" --upgrade $JET_INDEX_URLS + +# Keep this in the internal-only stage so public CI has no internal service dependency. +ARG CI_SERVER_URL +ARG NEMO_CI_TRIAGE_COMMIT=6e24e567ae2855e8acad5e2f780c7c26c195ab46 +RUN --mount=type=secret,id=NEMO_CI_TRIAGE_TOKEN \ + GIT_CONFIG_COUNT=1 \ + GIT_CONFIG_KEY_0=http.extraHeader \ + GIT_CONFIG_VALUE_0="Authorization: Basic $(printf 'oauth2:%s' "$(cat /run/secrets/NEMO_CI_TRIAGE_TOKEN)" | base64 -w0)" \ + uv pip install --no-cache-dir \ + "nemo-ci-triage @ git+${CI_SERVER_URL}/dl/nemo/nemo-ci-triage.git@${NEMO_CI_TRIAGE_COMMIT}" diff --git a/docker/common/install_nccl.sh b/docker/common/install_nccl.sh new file mode 100644 index 00000000000..786e6116e49 --- /dev/null +++ b/docker/common/install_nccl.sh @@ -0,0 +1,48 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +#!/bin/bash + +set -ex + +NCCL_VER="2.30.7-1+cuda13.3" + +for i in "$@"; do + case $i in + --NCCL_VER=?*) NCCL_VER="${i#*=}";; + *) ;; + esac + shift +done + +ARCH=$(uname -m) +if [ "$ARCH" = "amd64" ];then ARCH="x86_64";fi +if [ "$ARCH" = "aarch64" ];then ARCH="sbsa";fi + +curl -fsSLO https://developer.download.nvidia.com/compute/cuda/repos/ubuntu2404/${ARCH}/cuda-keyring_1.1-1_all.deb +dpkg -i cuda-keyring_1.1-1_all.deb +rm cuda-keyring_1.1-1_all.deb + +apt-get update + +if [[ $(apt list --installed | grep libnccl) ]]; then + apt-get remove --purge -y --allow-change-held-packages libnccl* +fi + +apt-get install -y --no-install-recommends \ + libnccl2=${NCCL_VER} \ + libnccl-dev=${NCCL_VER} \ + +apt-get clean +rm -rf /var/lib/apt/lists/* diff --git a/docs/api-guide/core/generalized_tensor_parallel.md b/docs/api-guide/core/generalized_tensor_parallel.md new file mode 100644 index 00000000000..dbba305e791 --- /dev/null +++ b/docs/api-guide/core/generalized_tensor_parallel.md @@ -0,0 +1,560 @@ +# Generalized Tensor Parallelism (GTP) + +> ⚠️ **Experimental.** GTP is an experimental feature and its API, configuration, and behavior may change in future versions without notice. + +> 📦 **Requires TransformerEngine >= 2.19** (GTP support is merged into TE main). On an older TE, GTP is disabled at import (`HAVE_GTP = False`) and enabling it raises an `ImportError` — please install TransformerEngine >= 2.19. + +**Generalized Tensor Parallelism (GTP)** is a lightweight, high-performance, memory-efficient distributed-training strategy implemented jointly in Megatron-LM and TransformerEngine. It **shards weight tensors across a GTP process group and reconstructs them on demand via asynchronous all-gather**, so larger models fit in the same memory without sacrificing throughput — the communication is overlapped with computation rather than added to it. + +GTP splits the weight-parallel domain into two orthogonal sub-axes — **`GTP = TP × GTP_remat`** — so every rank stores `1/(TP × GTP_remat)` of each linear weight, together with the matching slice of its gradient and optimizer state. + +**GTP_remat is an implementation of ZeRO-3**, and obeys the same contract: shard the weight (plus grad and optimizer state), all-gather it just before it is needed, use it, free it, reduce-scatter the gradient on the way back. What distinguishes it from the familiar ZeRO-3 / FSDP implementations is *where* it shards and *how finely* it materializes: + +- **It shards along a model-parallel axis, not the data-parallel one.** `GTP_remat` is a sub-axis of the weight-parallel grid that sits *on top of* TP — the `TP` slice stays sharded through the GEMM, and only the `GTP_remat` slice is rebuilt. It therefore composes with TP instead of competing with it for the same weight dimension. +- **It materializes one weight at a time, not a bucket.** Each `GTPShardedParam` gathers, computes and frees on its own schedule, which is what makes the per-weight prefetch chain (§3.4) and the low-precision gather (§1.3) possible — see the FSDP contrast in §1.1. + +| slice | stored | at GEMM time | +|---|---|---| +| **`TP`** | `1/TP` of the weight, permanently | **stays sharded** — ordinary tensor parallelism; the output is TP-sharded | +| **`GTP_remat`** | `1/GTP_remat` of the TP slice, permanently | **rematerialized**: all-gathered across the `GTP_remat` group just before the GEMM, so the GEMM sees the full TP slice; freed afterwards, and the wgrad is reduce-scattered back on the way out | + +Both `GTP_remat` collectives are prefetched one step ahead, so they overlap the previous layer's compute in forward *and* backward — the gather is off the critical path, not merely asynchronous. Note the two cuts do not always fall on the same axis of the weight (§1.4). + +**Turning it on.** The `GTP_remat` degree is `gtp_weight_remat_size`, derived from `--tensor-parallel-num-weight-shards` (= `tensor_model_parallel_size × gtp_weight_remat_size`). At **`gtp_weight_remat_size = 1` GTP is inactive and the path is byte-identical to plain TP + DP**, so it is safe to leave in the code path. It composes orthogonally with TP / SP / EP / DDP / CUDA Graphs. + +**Scope of this document**: a high-level summary of GTP_remat — design intent, public CLI surface, and Megatron-LM ↔ TransformerEngine integration touchpoints. + +**Source**: core implementation in `megatron/core/tensor_parallel/generalized_tensor_parallelism.py`, public surface re-exported from `megatron/core/tensor_parallel/gtp_api.py`. Low-precision tensor primitives (FP8 / MXFP8 / NVFP4) stay in TransformerEngine and are imported by the implementation module. + +**Outline:** + +- [Generalized Tensor Parallelism (GTP)](#generalized-tensor-parallelism-gtp) + - [1. Features](#1-features) + - [1.1 Fine-grained, per-weight materialization \& gradient reduction](#11-fine-grained-per-weight-materialization--gradient-reduction) + - [1.2 CUDA graph compatibility](#12-cuda-graph-compatibility) + - [1.3 Low-precision gather (native FP8 / NVFP4 param)](#13-low-precision-gather-native-fp8--nvfp4-param) + - [Per-microbatch schedule](#per-microbatch-schedule) + - [Communication volume breakdown](#communication-volume-breakdown) + - [GTP + NVFP4 (native NVFP4 param)](#gtp--nvfp4-native-nvfp4-param) + - [1.4 Composability with TP / SP / EP / DDP](#14-composability-with-tp--sp--ep--ddp) + - [1.5 Opt-in, minimally invasive integration](#15-opt-in-minimally-invasive-integration) + - [1.6 Optimizer-agnostic (Adam + Muon)](#16-optimizer-agnostic-adam--muon) + - [1.7 Scaling](#17-scaling) + - [1.8 Native distributed checkpointing (DCP)](#18-native-distributed-checkpointing-dcp) + - [2. Usage](#2-usage) + - [2.1 Required flags](#21-required-flags) + - [2.2 High-priority streams (Blackwell and later)](#22-high-priority-streams-blackwell-and-later) + - [2.3 Minimal end-to-end example](#23-minimal-end-to-end-example) + - [2.4 Tuning knobs](#24-tuning-knobs) + - [3. Implementation details](#3-implementation-details) + - [3.1 GTP\_remat architecture (Mcore ↔ TE integration)](#31-gtp_remat-architecture-mcore--te-integration) + - [What the flags do under the hood](#what-the-flags-do-under-the-hood) + - [Class hierarchy: which linears shard](#class-hierarchy-which-linears-shard) + - [Buffer / memory management](#buffer--memory-management) + - [Overlap design summary](#overlap-design-summary) + - [wgrad-before-dgrad schedule *(deferred to a follow-up MR)*](#wgrad-before-dgrad-schedule--deferred-to-a-follow-up-mr) + - [Recompute-forward prefetch chain *(GTP\_remat + activation recompute)*](#recompute-forward-prefetch-chain--gtp_remat--activation-recompute) + - [3.2 DDP buckets with (E)GTP\_remat](#32-ddp-buckets-with-egtp_remat) + - [3.3 Distributed checkpointing (DCP)](#33-distributed-checkpointing-dcp) + - [3.4 Prefetch-chain construction and its design assumptions](#34-prefetch-chain-construction-and-its-design-assumptions) + - [Grouped-expert chains (one-block-ahead)](#grouped-expert-chains-one-block-ahead) + - [4. Testing](#4-testing) + +--- + +## 1. Features + +### 1.1 Fine-grained, per-weight materialization & gradient reduction + +Each weight is sharded 1/N across a GTP_remat group along `out_features`, stored as a `GTPShardedParam` subclass of `nn.Parameter`. Materialization and gradient reduction are both **per-weight, per-call** — not per-model or per-module: + +- **Independent state per param**: each has its own AG state (`state`) and RS state (`rs_state`) machines, both cycling `NONE → ASYNC_WAIT → DATA_READY → NONE` and tracked separately so fwd and bwd async ops don't interfere. +- **Prefetch chain for AG** (doubly-linked `prev_w` / `next_w`): during fwd, each weight's `all_gather_and_prefetch` issues async AG for `next_w`; during bwd, `all_gather_and_prefetch_bwd` issues async AG for `prev_w`. Layer *i*'s AG overlaps with layer *i−1*'s GEMM. For an L-layer model, L−1 all-gathers are fully hidden behind compute. When activation recompute is enabled, a **third** chain prefetches the recompute-forward gathers during backward — see §3.1 *Recompute-forward prefetch chain*. One GEMM of runway covers a gather that stays inside the NVLink domain, but **not one that leaves it** — the case for MoE routed-expert weights, which also dominate the bytes gathered per block; those get their own *one-block-ahead* chains — see §3.4 *Grouped-expert chains*. +- **Deferred RS finalize for wgrad**: `wgrad_reduce_scatter` on param *i* launches an **async** reduce-scatter (handle stashed in `_wgrad_rs_handle`) and returns `None` to autograd — the wgrad is NOT finalized into `main_grad` yet. Finalization is **deferred one step**: the next bwd step (param *i−1*'s `wgrad_reduce_scatter`) calls `self.next_w._wait_reduce_scatter()` + `_finalize_wgrad()`, which waits on the stashed handle, accumulates the reduced wgrad into `main_grad`, and fires the DDP `register_grad_ready` hook. The chain's head (first-in-fwd, last-in-bwd) uses a synchronous RS since nothing follows it. This one-step deferral is what lets layer *i*'s RS overlap with layer *i−1*'s bwd GEMMs. +- **Cold start only**: every weight's very first AG is synchronous (`DATA_READY_SYNC`, no prefetch has run yet); the async prefetch chain kicks in from the second forward onward. + +Contrast with FSDP: FSDP gathers at module-group granularity in full precision with PyTorch-managed lifecycle. GTP_remat works at individual-weight granularity, in quantized form, with its own explicit ticket-based buffer pool and a one-step-deferred RS finalizer. + +> **FSDP can't shrink into GTP_remat because FSDP's overlap is bucket-grained by design** — bucket granularity exists *to avoid* paying NCCL launch latency on tiny params (LayerNorm γ/β, biases, Mamba `dt_bias`/`D`/`A_log`) and *to avoid* the per-weight scheduling state that GTP_remat relies on (per-param prefetch chain, ticket-based buffer cache, stream choreography). Removing buckets doesn't make FSDP faster; it makes FSDP into GTP_remat, with all the engineering that entails — selective wrapping (only large GEMM weights), per-weight prefetch chain, per-param buffer ticket, and explicit AG/RS stream choreography on a side stream so external drains have something meaningful to wait on. + +### 1.2 CUDA graph compatibility + +CG compatibility is designed-in from day one, not retrofitted. The entire sync / buffer / chain architecture is shaped around making **captured fwd/bwd replays produce identical bit-for-bit behavior** — without the usual capture-vs-eager pitfalls that force other weight-sharding schemes to either disable CG or require special handling. + +- **Chains never cross-link across the capture axis** (`GTPChain.GRAPHED` / `GTPChain.UNGRAPHED`, plus the eager-only grouped-expert chains of §3.4). `prev_w` / `next_w` only connect same-chain params, so a captured traversal never reaches into eager Python and vice-versa. +- **`torch.cuda.Event(external=True)`** for `ag_event` / `rs_event` — the events survive CG capture boundaries and can be waited on from replay-time streams. +- **Idempotent ticket cache**: `GTPWeightCache.get(ticket)` keeps `slot.buf` set even after `release()`, so replays read the same buffer address as capture. `clear()` drops buffers while keeping tickets valid → supports CG re-capture with lazy re-allocation. +- **Allocate-in-pool at creation** (`set_cuda_graph_mempool` + `_graphed_alloc`): GRAPHED-chain AG/RS buffers and quantized weight storage are allocated **directly into the CG memory pool** at first creation (during warmup, before capture), so no CUDA allocations happen inside the captured graph — and no post-hoc reallocation/clone is needed. UNGRAPHED buffers stay in regular allocator memory. +- **Lazy, one-shot chain linking**: `prefetch_initialized` is flipped during the first fwd (warmup), so the chain-construction Python side-effects never execute inside a captured graph. The link table is buffered and flushed atomically at the second forward. +- **DDP hook manual triggering**: `register_grad_accum_hook` stores the DDP hook on the param; `_CudagraphReplayNode.backward` calls it manually after replay (since `AccumulateGrad` hooks are silenced by replay). This is also how the `assert self.grad_reduce_handle is not None` failure from partial-CG + overlap-grad-reduce is resolved. +- **Warmup is side-effect-free on `main_grad`**: GTP_remat accumulates wgrad into `main_grad` *inside* the backward (the fusion path returns wgrads as graph outputs instead). Graph capture only *records* ops; it never runs them. But `create_fwd_graph` runs an **eager** warmup fwd+bwd before capturing. That warmup backward executes GTP_remat's `main_grad.add_`. Its deferred cascade adds into a cross-graph `next_w` (another module) from a **stale RS ticket** — the prior backward's wgrad. And `create_cudagraphs()` runs *after* `finalize_model_grads`. So this overwrites the finalized (reduced + per-token-scaled) grads and spikes the step's grad norm. **Fix**: `create_fwd_graph` snapshots the grads its warmup touches — own params + cross-graph `next_w` — via `_backup_grads_before_capture`, then restores them after capture. The bwd graph has no warmup, so it needs none. Bounded to one module's grads. +- **Drains at CG / eager boundary**: `_drain_gtp_side_streams()` before eager MoE expert compute. Inside bwd capture, two-phase drain: Phase 1 joins the within-graph cascade and records `bwd_completion_event` (next runner unblocks); Phase 2 calls `wait_async_comms(GRAPHED)` to drain the chain-tail handle and re-joins side streams (queued after the event so it doesn't delay the next runner). +- **Side-stream registration**: the `(GRAPHED, gtp_remat_group)` ag/rs streams are materialized at runner init (`_register_gtp_side_streams`) so they are captured before the first forward. + +### 1.3 Low-precision gather (native FP8 / NVFP4 param) + +Wire bandwidth scales with the **quantized** size, not BF16 size — GTP_remat composes with low-precision training rather than fighting it. The shard is stored as a native **MXFP8**, native **NVFP4**, or **BF16** weight, gathered with the following mechanics: + +- **Native MXFP8 param — `mxfp8` + `--fp8-param-gather` (always paired, see §2.1).** The shard **is** a native `MXFP8Tensor` (§3.1); the optimizer writes FP32 master → FP8 once per step (off the forward critical path), and the forward **all-gathers the FP8 shard directly** — no per-microbatch quantize, no cast. The rowwise (fwd) / columnwise (bwd) view comes from a *separate* gather-quantizer copy (`_gtp_gather_quantizer`), leaving the param's own quantizer for the optimizer's write path. +- **Native NVFP4 param — `--fp4-param-gather` (required).** Same shape as MXFP8: the shard **is** a native `NVFP4Tensor`, all-gathered as packed 4-bit (`kFloat4E2M1`) and optimizer-maintained, no per-microbatch quantize. See the *GTP + NVFP4* subsection below. +- **BF16 (no FP8/NVFP4 params).** The BF16 shard is all-gathered as-is. +- **Coalesced NCCL**: `grouped_gather_along_first_dim` uses `torch.distributed._coalescing_manager` to batch E experts' AGs into a single NCCL op. +- **Padding**: shards are allocated **already padded** so each rank's dim0 stays `pad_for_alignment`-divisible (MXFP8: 32). Column-parallel pads the per-TP slice (`out_features / tp_size`) to a multiple of `pad_for_alignment × gtp_remat_size` so it survives TE's TP split aligned; row-parallel / Megatron-local pad the TP-local tensor directly (§3.1). Padding lands contiguous at the tail, so stripping is one trailing slice (`tensor[:-pad_length]`). + +#### Per-microbatch schedule + +``` +Steady-state fwd (MXFP8 native FP8 param / BF16): + default: ──GEMM(W_0)───────────────────GEMM(W_1)───────────────────GEMM(W_2)──... + ag_str: [AG_issue W_1] [AG_issue W_2] + (no per-microbatch quantize: the FP8 shard is + maintained by the optimizer; BF16 gathers as-is) + +Steady-state bwd (MXFP8 / BF16): + default: ──bwd GEMMs(W_i)──... + ag_str: [AG_issue W_{i-1}] + (columnwise view of the same FP8 shard; no quant) +``` + +For the native-FP8 (MXFP8), native-NVFP4, and BF16 paths the forward all-gather is a **single** NCCL op per weight on the GTP_remat ncclStream, with no per-microbatch quantize or GTP_remat-group amax on the critical path (the standard DP-group FP8 amax allreduce in `reduce_and_update_fp8_tensors` is unchanged by GTP_remat). Only the `dist.all_gather` issue is wrapped in `with torch.cuda.stream(ag_stream)`; the NCCL kernel runs on c10d's private ncclStream and overlaps with the next GEMM until it reaches its wait. + +#### Communication volume breakdown + +Per-microbatch per-weight comm budget (assuming bf16 wgrad reduce-scatter): + +| Format | Block | Data B/elem | Scale_inv B/elem | Per-elem | Fwd AR(amax) | Fwd AG | Bwd AG | Wgrad RS (bf16) | Total B/elem | vs BF16 | +|--------|-------|-------------|------------------|----------|--------------------------------|--------|--------|-----------------|--------------|----------------| +| BF16 | n/a | 2.0000 | — | 2.0000 | — | 2.0000 | 2.0000 | 2.0000 | 6.0000 | 1.00× (baseline) | +| MXFP8 | 32 | 1.0000 | 1/32 = 0.0313 | 1.0313 | — (microscale, no global amax) | 1.0313 | 1.0313 | 2.0000 | 4.0626 | 0.68× (–32%) | +| NVFP4 | 16 | 0.5000 | 1/16 = 0.0625 | 0.5625 | — (scale set at opt-step quantize) | 0.5625 | 0.5625 | 2.0000 | 3.1250 | 0.52× (–48%) | + +How to read the columns: +- `Per-elem` = `Data B/elem + Scale_inv B/elem` — wire cost of one quantized weight buffer (data + scale_inv together). +- `Fwd AG` and `Bwd AG` each carry the quantized buffer once, so they equal `Per-elem`. Bwd all-gathers the same FP8 shard (columnwise view) — no re-quantize, no AR(amax). +- `Wgrad RS (bf16)` = 2.0 B/elem — gradient is reduce-scattered in bf16 regardless of weight precision. +- `Fwd AR(amax)` — none per microbatch for either native format: MXFP8 is microscale-only, and native NVFP4 carries its block scales in the gathered buffer with the per-tensor scale set at the optimizer-step quantize (not per forward). +- `Total B/elem` = `Fwd AG + Bwd AG + Wgrad RS` — there is no per-microbatch amax AR to add. + +Gathering the pre-quantized weight attacks AG only: the AG portion shrinks ~72% from BF16 → NVFP4, but RS is untouched, so the wgrad RS becomes the dominant comm path in NVFP4 (~64% of the budget at bf16 RS, ~78% at fp32 RS). + +#### GTP + NVFP4 (native NVFP4 param) + +NVFP4 GTP_remat keeps each shard as a native `NVFP4Tensor` and all-gathers it as packed 4-bit (`kFloat4E2M1`) — the native-param path, mirroring native MXFP8: the distributed optimizer writes the NVFP4 shard directly once per step and the forward all-gathers it with no per-microbatch quantize. + +- **`--fp4-param-gather` is mandatory.** Without it NVFP4 GTP falls back to a BF16 all-gather that trips TE's scaling-mode assert (`DELAYED` vs `NVFP4`); `validate_args` enforces it and raises early. +- **Mixed-precision models (per-layer quant config).** A model may assign recipes per layer — e.g. NVFP4 default, MXFP8 for `mixer.out_proj`, BF16 for attention (`linear_qkv`/`linear_proj`) and latent MLPs. NVFP4 params gather natively as above. **MXFP8 params cannot be native-param-gathered** — the DDP param buffer has no MXFP8 storage remap (`replace_raw_data` is unimplemented for `MXFP8Tensor`, unlike NVFP4's packed-rowwise remap), so they are all-gathered in **BF16** and re-quantized with the layer's **own MXFP8 quantizer** inside the TE backward dgrad path (not the global delayed recipe). BF16-recipe layers gather BF16 unchanged. + +### 1.4 Composability with TP / SP / EP / DDP + +- **TP** (intra-layer): orthogonal axis — GTP_remat shards `out_features` regardless of TP's parallel mode (column or row). 2D grid naturally formed via `tp_group × gtp_remat_group`. + +> ⚠️ **The two cuts are not always on the same axis.** `GTP_remat` **always** slices `out_features` (dim 0) of the TP-local weight — independent of TP's `partition_dim`: +> +> | linear | TP cuts | `GTP_remat` cuts | | +> |---|---|---|---| +> | **column-parallel** (`linear_qkv`, `linear_fc1`) | `out_features` | `out_features` | same axis → `out_features/(TP × GTP_remat)` | +> | **row-parallel** (`linear_proj`, `linear_fc2`) | `in_features` | `out_features` | **perpendicular** → `in_features/TP` × `out_features/GTP_remat` | + +- **SP** (sequence-parallel): transparent — GTP_remat operates at weight dim, SP at sequence dim. +- **EP** (MoE): `GroupedLinear` with GTP_remat → each routed expert sharded across `EXPERT_GTP_WEIGHT_REMAT_GROUP`, independent of EP. MoE AllToAll (HybridEP/NVLink) runs independently of GTP_remat AG/RS (NCCL/IB). +- **DDP**: GTP_remat bypasses autograd's grad accumulator (async RS returns `None`; `_finalize_wgrad` accumulates directly into `main_grad`). DDP registers its grad-ready hook on GTP_remat params via `register_grad_accum_hook` (not autograd's `AccumulateGrad`); GTP_remat invokes it from `_finalize_wgrad` (eager path) and `_CudagraphReplayNode.backward` (captured path) **after** the wgrad lands in `main_grad`, so a bucket's DDP reduce-scatter runs strictly after every GTP_remat param's `{RS → main_grad add}` — never over a stale `main_grad` — and DDP↔GTP_remat NIC deadlock at IB scale is avoided. See §3.2. + +### 1.5 Opt-in, minimally invasive integration + +- **TE is GTP-agnostic.** Mcore builds the plain TE linear with an already-sharded `out_features` and attaches a `GTPShardedParam` *after* construction; TE dispatches through its generic **`DistributedWeight` protocol** (gates on `is_distributed_weight`) and takes no GTP argument, so there is no framework-level refactor and callers never thread a group (§3.1). +- **Opt-in by linear *class*; sharding stays per-*weight*.** *Which* linears opt in is class-based — GTP_remat wraps the TE classes that resolve a shard group internally (`TEColumnParallelLinear` / `TERowParallelLinear` / `TELayerNormColumnParallelLinear` for dense, `TEGroupedLinear` for routed experts), so upper-level modules thread no `gtp_remat_group`. But materialization and gradient reduction stay at **individual-weight** granularity — each wrapped weight is its own `GTPShardedParam`, gathered/reduce-scattered per-weight, per-call (§1.1). Base `TELinear` (e.g. MoE latent-proj MLPs) and small replicated tensors (LayerNorm γ/β, biases, Mamba `dt_bias`/`A_log`/`D`/`conv1d`, MoE router) **stay full** — the all-gather wouldn't amortize (§3.2 *dense non-GTP_remat* vs *dense GTP_remat*). +- **Off is a byte-for-byte no-op.** When the resolved group is `None`/size-1, `_gtp_pre_init` leaves `out_features` unsharded and `_gtp_attach_post_init` short-circuits (as does `wrap_module_params_gtp` for Megatron-local linears); when `gtp_weight_remat_size == 1` the `layers.py` GTP_remat path is skipped entirely. +- **Chain setup is one pass.** `classify_gtp_chains(model)` walks `named_parameters()` once at init and sets `chain_id` on every `GTPShardedParam` from the current `cuda_graph_modules` (§3.4). +- **Knobs.** `GTPRematConfig.{pad_for_alignment, weight_prefetch, check_param_states}`, plus the debug-name tagger `tag_gtp_params_with_names` for readable link-table output. + +### 1.6 Optimizer-agnostic (Adam + Muon) + +GTP_remat runs under both the standard **Adam** `DistributedOptimizer` and **Muon** (the `LayerWiseDistributedOptimizer`), DCP save/load included: + +- **Adam** shards optimizer state over the gtp_remat/egtp_remat-excluded replicate group, like any GTP_remat run (§3.2). +- **Muon** keeps matrix params *whole* (Newton–Schulz needs the full 2D weight). A GTP_remat-replicated whole param (e.g. MoE router, latent-proj MLPs) then lands on one checkpoint key shared by all GTP_remat peers, so the LayerWise optimizer folds `gtp_rank` into its `replica_id` — exactly one peer writes (the optimizer-state analog of the model-side fold in §3.3). +- **Native-FP8 optimizer-state matching (Muon path).** The save-side dequantize (§3.3) hands DCP a *fresh* BF16 tensor, which breaks the id-based optimizer-param → model-`ShardedTensor` match for every native-FP8 GTP_remat weight. The dequantized copy carries a `_gtp_dequant_src` backlink to the live FP8 param, and `_backfill_gtp_sharded_param_map` reuses the model's **own** entry (backlink first, tagged-name second) — preserving its full offsets (expert axes included) and `replica_id`. Only truly-unmatched params (Mamba `in_proj`, a gathered+split factory) take the per-shard rebuild, which refuses expert-parallel params rather than emit EP-colliding shards. + +Neither path adds a GTP_remat-specific checkpoint format or call site. + +### 1.7 Scaling + +Effective per-GPU weight size = `W / (TP × GTP_remat)`. Example: TP=4 + GTP_remat=8 with NVFP4 → 32× weight-memory reduction and 128× wire-bandwidth reduction vs full BF16 replication, before data parallelism. + +**Weak scaling.** GTP_remat fixes the shard width and grows the job by adding data-parallel replicas (DP = #GPUs / GTP_remat), so per-GPU compute stays constant while only the DP gradient reduction widens with scale. + +The best GTP_remat size is model- and cluster-dependent — driven by weight sizes, per-GPU memory headroom, and which collectives can be kept on fast links — so there is no single recommended value. The example below runs on **GB200 NVL72** (a 72-GPU NVLink domain) and uses **GTP64**, which places communication as: + +- **NVLink-local:** the *dense-layer* (Mamba / attention / shared-expert) GTP_remat weight all-gather + wgrad reduce-scatter, **and** the `EP64` all-to-all dispatch/combine — all kept inside one ≤72-GPU NVLink domain (EP64 ≤ NVL72). +- **Inter-node (IB / CX7):** the DP gradient reduction **plus** the `EGTP2` expert-weight all-gather / wgrad reduce-scatter, whose 2 shards land on different NVLink domains and so cross nodes. + +On an Ultra-proxy hybrid Mamba-MoE model (**~280B parameters**; `GTP64 · EP64 · EGTP2`, mb1, MXFP8, BF16 reduce-scatter, no CUDA graph), scaling efficiency holds **≥93 % of the single-domain (128-GPU / DP2) baseline out to 3072 GPUs (DP48)**, while max reserved memory *decreases* with scale (137 → 104 GB) as the distributed optimizer shards optimizer/grad state across more DP replicas. + +> **Takeaway:** near-flat weak scaling — **≥93 % efficiency from 128 → 3072 GPUs**, with per-GPU memory shrinking as DP grows. + +![GTP64 weak-scaling efficiency](../../images/generalized_tensor_parallel/0617_gtp64_weak_scaling_efficiency.png) + +### 1.8 Native distributed checkpointing (DCP) + +**GTP_remat + DCP is straightforward:** +- Reuses the existing checkpoint stack rather than adding a parallel one. GTP_remat-sharded weights *and* distributed-optimizer state save/load through the standard PyTorch / Mcore `torch_dist` sharded checkpoint, with **no GTP_remat-specific format or call path** and a tiny code footprint (one new helper + one helper made GTP_remat-aware). +- Checkpoints **reshard freely** across different `(TP, GTP_remat, EGTP_remat, DP, PP)` topologies — including a different GTP_remat/EGTP_remat size — with no offline conversion. + +See [§3.3 Distributed checkpointing (DCP)](#33-distributed-checkpointing-dcp) for details. + +--- + +## 2. Usage + +GTP_remat is enabled through two CLI flags on Megatron's training launcher; everything else (process-group construction, parameter slicing, prefetch chain wiring, optimizer routing) is automatic once the flags are set. + +### 2.1 Required flags + +```bash +# Total number of shards each dense weight (attention, mamba, MLP linears) is split into along +# out_features, across the tensor-parallel + GTP_remat axes. Must be >= --tensor-model-parallel-size and +# divisible by it. The GTP_remat degree is derived as num_weight_shards / tensor_model_parallel_size +# (e.g. TP=1 + num_weight_shards=2 -> GTP_remat=2; TP=2 + num_weight_shards=8 -> GTP_remat=4). +--tensor-parallel-num-weight-shards + +# Total number of shards each MoE routed-expert weight is split into along out_features, across the +# expert-tensor-parallel + expert-GTP_remat axes. Must be >= --expert-tensor-parallel-size and divisible +# by it. The expert-GTP_remat degree is derived as num_weight_shards / expert_tensor_parallel_size. +# Independent from --tensor-parallel-num-weight-shards; can be left unset for non-MoE models. +--expert-tensor-parallel-num-weight-shards +``` + +> The (dense / expert) GTP_remat degree is exposed **only** through +> `--tensor-parallel-num-weight-shards` / `--expert-tensor-parallel-num-weight-shards`. The internal +> `gtp_weight_remat_size` / `expert_gtp_weight_remat_size` config fields are derived from them and +> have no CLI flag. + +**Low precision (MXFP8).** GTP_remat + `--fp8-recipe mxfp8` **requires** both `--fp8-param-gather` +and `--reuse-grad-buf-for-mxfp8-param-ag` (`arguments.py` asserts this) — the weight is a native FP8 +param, and since MXFP8 cannot map into the contiguous param buffer (`replace_raw_data` unsupported) +the all-gather reuses the grad buffer. Mechanism: §1.3, §3.1. + +**Low precision (NVFP4).** GTP_remat + `--fp4-format` **requires** `--fp4-param-gather` +(`arguments.py` asserts this) — without it NVFP4 weights fall back to a BF16 gather that fails the +backward GEMM. Mechanism and mixed-recipe (MXFP8-override) handling: §1.3 → *GTP + NVFP4*. + +### 2.2 High-priority streams (Blackwell and later) + +Required on GB200 / GB300 so the GTP_remat comm streams get the SM priority needed for AG/RS overlap with compute: + +```bash +--high-priority-stream-groups ep gtp_remat expt_gtp_remat tp +``` + +The launcher also exports `CUDA_GRAPHS_USE_NODE_PRIORITY=1` so captured CUDA graphs respect the inherited stream priority. + +### 2.3 Minimal end-to-end example + +```bash +# 4 ranks, TP=2 + GTP_remat=2 across out_features, BF16 weights. +# TP=2 + num-weight-shards=4 -> GTP_remat = 4 / 2 = 2. +torchrun --nproc-per-node 4 pretrain_gpt.py \ + --tensor-model-parallel-size 2 \ + --pipeline-model-parallel-size 1 \ + --tensor-parallel-num-weight-shards 4 \ + --expert-tensor-parallel-num-weight-shards 1 \ + --high-priority-stream-groups ep gtp_remat expt_gtp_remat \ + --bf16 \ + --num-layers 12 --hidden-size 1024 --num-attention-heads 16 \ + --seq-length 1024 --max-position-embeddings 1024 \ + --micro-batch-size 1 --global-batch-size 4 \ + --train-iters 10 \ + --use-mcore-models \ + --transformer-impl transformer_engine \ + --tokenizer-type NullTokenizer --vocab-size 32000 \ + --data-path --split 99,1,0 +``` + +At iter-0 you'll see one rank-0 log line confirming the active config: + +``` +GTP_remat enabled. GTPRematConfig(pad_for_alignment=16, check_param_states=False, + weight_prefetch=True, async_reduction=True, calculate_per_token_loss=False) +``` + +### 2.4 Tuning knobs + +Set via `from megatron.core.tensor_parallel.generalized_tensor_parallelism import GTP_CONFIG, update_gtp_config`: + +```python +update_gtp_config( + pad_for_alignment=16, # NVFP4: 16, MXFP8: 32, BF16: any; auto-set in training.py + weight_prefetch=True, # Disable to debug the cold-start path + async_reduction=True, # Whether to perform GTP_remat gradient reduction asynchronously + calculate_per_token_loss=False, # Mirror config.calculate_per_token_loss (SUM vs MEAN RS) +) +``` + +`training.py` auto-tunes `pad_for_alignment` based on the quantization recipe (`--fp4`, `--fp8-recipe=mxfp8`, etc.) before model construction. The other knobs are usually left at defaults. + +> **CUDA-graph warmup under GTP_remat.** When CUDA graphs are enabled, GTP_remat forces a minimum of **2** per-graph warmup steps regardless of `--cuda-graph-warmup-steps` (e.g. a user-set `0` is bumped to `2`): the first warmup builds the weight-prefetch chain and the second exercises the prefetch path before capture. + +--- + +## 3. Implementation details + +### 3.1 GTP_remat architecture (Mcore ↔ TE integration) + +![GTP_remat / Mcore-TE integration architecture](../../images/generalized_tensor_parallel/0712_gtp_te_protocol_redesign.png) + +**Ownership.** TE owns the linear primitives (`Linear` / `LayerNormLinear` / `LayerNormMLP` / `GroupedLinear`), the low-precision tensor types (FP8 / MXFP8 / NVFP4), and a generic **`DistributedWeight` protocol** (`transformer_engine/pytorch/distributed_weight.py`). Megatron owns **all** GTP_remat logic — sharding, the prefetch chain, the buffer cache, the AG/RS state machines, and DDP integration. **TE never names GTP.** + +**The bridge** — three touch points, nothing more: + +- **Construction.** Mcore pre-shards `out_features` (`_gtp_pre_init`) so plain TE builds *this rank's shard* directly; GTP is attached *after* build (`_gtp_attach_post_init`). TE takes no GTP argument. +- **Runtime.** TE's fwd/bwd gate on `is_distributed_weight(weight)` and call the generic list-shaped dispatchers (`materialize_weight_for_forward` / `materialize_weight_for_backward`, `finalize_weight_grads`). `GTPShardedParam` implements the protocol (`materialize_group_for_forward`/`_backward`, `finalize_group_grads`, `grad_buffer`); the concrete collectives (`all_gather_and_prefetch`, `wgrad_reduce_scatter`) live only in Megatron. A plain tensor is a no-op. +- **Streams.** `_register_gtp_side_streams` / drain calls synchronize TE's GEMMs with the side stream that owns the AG/RS NCCL ops. + +**One init path, all precisions.** Since `out_features` is pre-sharded, TE builds the shard directly — native `MXFP8Tensor` (`--fp8-param-gather`), native `NVFP4Tensor` (`--fp4-param-gather`), or BF16 — **with no full weight ever materialized**. `attach_gtp_to_presharded_module` then turns it into a `GTPShardedParam`: a native quantized shard is reclassed in place to `GTP_` (stays buffer-resident on the quantized dist-opt path); a BF16 shard is re-registered (no slice — already shard-sized). The optimizer maintains the shard end-to-end, gathered each forward with **no per-microbatch re-quantize** (§1.3). + +> **Per-GTP-rank init.** Each rank draws its *own* shard, so GTP weights need *distinct* random values per GTP_remat peer (else the gather would be `gtp_remat_size` identical blocks). `model_parallel_cuda_manual_seed` adds `gtp-remat-rng` / `egtp-remat-rng` trackers (offset per peer) that `_gtp_pre_init` routes init through; replicated params keep the shared trackers. Added only when the axis is active, so non-GTP runs keep a byte-identical tracker set. + +> **Megatron-local linears** (`ColumnParallelLinear` etc. in `tensor_parallel/layers.py`) still build the full weight and slice post-init via `wrap_module_params_gtp` — unchanged. + +#### What the flags do under the hood + +The `--*-num-weight-shards` flags flow through five stages, from process groups to the prefetch chain: + +1. **Process groups.** `initialize_model_parallel(...)` treats GTP_remat/EGTP_remat as **first-class orthogonal axes** (`world = TP·GTP_remat·CP·DP`; experts `= ETP·EP·EGTP_remat·PP·expert_dp`), building `_GTP_WEIGHT_REMAT_GROUP` and `_EXPERT_GTP_WEIGHT_REMAT_GROUP` (sizes = `num-weight-shards / TP` and `/ ETP`). **DP and gtp_remat stay orthogonal:** `get_data_parallel_group()` is the replicate axis (DDP + optimizer shard over it); `with_gtp_remat=True` gives the combined DP × gtp_remat axis for data distribution. + + > **Batch-size arithmetic.** `args.data_parallel_size` is the **replicate degree only** — gtp_remat is *divided out* of it (folded into `total_model_size` at `arguments.py:446`). But data is distributed over the **full DP × gtp_remat axis**, so each gtp_remat peer consumes a *distinct* microbatch and the global sample count is `micro_batch_size × data_parallel_size × gtp_weight_remat_size × num_microbatches`. The training loop therefore **re-applies `gtp_weight_remat_size`** to close the gap: *multiplied back in* for the LR-scheduler `increment` and the logged `batch_size`, *divided back out* to recover `eval_num_microbatches`. Without this it would read as a double-count — it is not. + +2. **Per-class sharding.** `extensions/transformer_engine.py` decides *per linear class* whether to shard, so **no `gtp_remat_group` is threaded through the module APIs** (attention, Mamba, MLP, embedding, MTP). Dense wrappers resolve the group via `utils.get_gtp_weight_remat_group(...)`; `TEGroupedLinear` uses `pg_collection.expt_gtp_remat`. Group `None`/size-1 → left full; otherwise `_gtp_pre_init` pre-shards `out_features` and `_gtp_attach_post_init` makes the shard a **`GTPShardedParam`** (the `DistributedWeight` implementer; native FP8/NVFP4 by reclass, BF16 by re-register). Base `te.Linear` (MoE latent projections) gets no group and stays full → see [Class hierarchy](#class-hierarchy-which-linears-shard). + +3. **Gradients (DDP).** GTP_remat shards are ordinary DDP params in the usual dense/expert buffers, reduced over the **replicate** group. The gtp_remat axis is completed separately: **GTP shards by their reduce-scatter, replicated params by an all-reduce** in `finalize_model_grads` (mean-vs-sum per `calculate_per_token_loss`) → see §3.2. + +4. **Optimizer.** State is sharded over the same replicate group; **global-norm clipping** reduces over the dist-opt grad-stats group spanning the full world (incl. gtp_remat/egtp_remat), counting replicated params **once per axis** to avoid over-counting. + +5. **Prefetch chains.** `classify_gtp_chains(model)` runs once after build (`get_model`) and wires each `GTPShardedParam` into a **`GRAPHED`/`UNGRAPHED`** chain from `cuda_graph_modules` → see [§3.4 Prefetch-chain construction](#34-prefetch-chain-construction-and-its-design-assumptions). + +#### Class hierarchy: which linears shard + +The figure visualizes the per-class split from the list above: green = resolves a GTP_remat group and shards, red = base `TELinear` (MoE latent projections) that stays full. Dashed arrows are *builds* (module → leaf); solid arrows are *inherits* (leaf → TE primitive). + +![GTP_remat class hierarchy — which TE linear classes shard](../../images/generalized_tensor_parallel/0628_gtp_remat_class_hierarchy.png) + +#### Buffer / memory management + +Two distinct pools with explicit lifecycle rules: + +- **`GTPWeightCache`** (AG/RS output buffers) — ticket-based, keyed on `(shape, dtype, fwd, expert_idx, reduce_scatter)`. Same-shape buffers across layers are shared. Tickets persistent; buffer allocated lazily on first `get()`; addresses stable across iterations for CG replay. +- **`_wgrad_buf_pool`** (wgrad-GEMM output recycling) — holds the **full, unsharded** wgrad-GEMM output buffer (shape `_unsharded_shape`, dtype `main_grad.dtype` — fp32 when `grad_reduce_in_fp32`, else bf16). The TE backward writes the wgrad into it via `main_grad_func = weight.grad_buffer` (a `DistributedWeight` protocol method backed by `get_wgrad_tensor`; it is a *scratch*, distinct from the sharded `param.main_grad`); the protocol's `finalize_group_grads` (backed by `wgrad_reduce_scatter`) then reduce-scatters it down to the shard and the buffer is returned here. This is a full-weight-shaped fp32/bf16 transient — one of the larger per-weight buffers — and is **precision-independent** (wgrad is always computed in high precision), so it is identical in BF16 vs MXFP8 runs. Buffers are tagged `_from_gtp_wgrad_pool=True` at `_wgrad_pool_get`; `_wgrad_pool_put` no-ops on foreign buffers (fresh allocs from Megatron `layers.py` or aten F.embedding bwd) → caching allocator handles those, so the pool never accumulates untagged buffers. + +#### Overlap design summary + +``` +fwd: AG(W_{i+1}) ∥ GEMM(W_i) ∥ CG replay of captured layers +bwd: AG(W_{i-1}) ∥ dgrad(W_i) → wgrad(W_i) ∥ RS(wgrad_i) ∥ [finalize wgrad_{i+1} + DDP hook] +``` + +GTP_remat runs up to **three** independent prefetch chains, all following one rule — *prefetch the weight the next consume will need*: + +| # | when | consume | prefetch (overlap) | AG direction | slot | +|---|------|---------|--------------------|--------------|------| +| 1 | fwd | weight `i` | `next_w` = i+1 ‖ `GEMM_i` | rowwise (`fwd=True`) | `_prefetch_handle` | +| 2 | bwd dgrad | weight `i` | `prev_w` = i−1 ‖ `Dgrad_i` | columnwise (`fwd=False`) | `_prefetch_handle` | +| 3 | bwd recompute | weight `i` | `_recompute_next` = i+1 ‖ `recompute_GEMM_i` | rowwise (`fwd=True`) | `_recompute_prefetch_handle` (separate) | +| 1b | fwd (MoE, eager) | expert weight `i` | same role in MoE block i+1 ‖ *whole block i* | rowwise (`fwd=True`) | `_prefetch_handle` | + +Row 1b is chain 1 applied to a *homogeneous* chain: routed-expert `fc1`/`fc2` link across consecutive MoE blocks, so the runway is a full block rather than one GEMM (§3.4 *Grouped-expert chains*). + +Chain 3 exists only when activation recompute is on. It mirrors chain 1 (rowwise, prefetch `next`) but runs *during* backward, so it overlaps chain 2 in time on the same weight — hence its **own** slot. fwd (1) and bwd-dgrad (2) never overlap in time, so they safely share `_prefetch_handle`. See *Recompute-forward prefetch chain* below. + +At bwd step *i* the step is launching *RS of wgrad_i* while finalizing the *previous* iter's wgrad (`wgrad_{i+1}` in bwd order = the next-one-over in fwd order). That one-step deferral is what makes the RS run concurrent with the next layer's dgrad/wgrad GEMMs instead of blocking after every layer. + +Communication never blocks compute except at the very first layer of each direction (cold start) and at enforced serialization points (CG/eager drains, finalize-grads barrier). + +##### wgrad-before-dgrad schedule *(deferred to a follow-up MR)* + +Current behavior: backward always runs dgrad GEMM, then wgrad GEMM, then issues the GTP_remat wgrad RS — the RS overlaps with the *next* layer's bwd GEMMs (the one-step deferral above). + +A future MR will add an opt-in wgrad-before-dgrad schedule on `_Linear` / `_LayerNormLinear` so the GTP_remat wgrad RS NCCL overlaps with the dgrad GEMM of the **same** layer (best for the GTP_remat + no-TP case). + +##### Recompute-forward prefetch chain *(GTP_remat + activation recompute)* + +When a GTP_remat-sharded module is in `--recompute-modules` (e.g. `shared_experts`), its forward is **re-run during backward** to regenerate activations. That recompute-forward must all-gather each weight **rowwise** again — a *third* gather lifecycle, concurrent with the in-flight **columnwise** dgrad gather of the *same* weight. Since both share one `GTPShardedParam`, the recompute path gets its **own** prefetch slot (`_recompute_prefetch_handle` / `_recompute_ag_event`, reusing the `_ag_ticket_fwd` rowwise buffer) so it never clobbers the dgrad lifecycle's `state` / `_prefetch_handle` / `ag_event`. + +The recompute weights form a **separate** linked list (`_recompute_next`), **self-populated** on the first backward from the weights actually re-gathered while `in_fp8_activation_recompute_phase()` is true — membership is *observed*, not configured (no tagging, so it tracks exactly what each checkpointed module re-gathers). Each recompute-forward consume prefetches the next recompute weight, so every gather **except the global-first** overlaps preceding recompute / dgrad / wgrad compute: + +``` +recompute-fwd of shared_experts (per layer: GEMM fc1 → SReLU → GEMM fc2, then dgrad+wgrad) + + Before (on-demand): + default: AG(fc1)─GEMM fc1─SReLU─AG(fc2)─GEMM fc2─dgrad─wgrad─... every AG exposed + After (recompute chain): + default: GEMM fc1─SReLU─GEMM fc2─dgrad─wgrad─GEMM fc1'─... back-to-back + ag_str: AG(fc1) [AG fc2] [AG fc1' (next layer)] only AG(fc1) exposed +``` + +`AG(fc2)` is issued at `fc1`'s consume (overlaps GEMM fc1 + SReLU); `AG(fc1')` for the next layer is issued at `fc2`'s consume, so it overlaps the **whole** layer's `dgrad + wgrad` window. The cross-layer link is what hides every region head except the very first. + +Under **full-iteration CUDA graphs** the recompute-forward is captured; `wait_async_comms(GRAPHED)` drains the recompute handle too (sets `_recompute_already_drained`) so the captured consumer skips its cross-graph wait — the same producer-drain pattern as the fwd/bwd chains. + +> **When *not* to recompute a GTP_remat weight.** Recompute on a GTP_remat-sharded weight adds this extra rowwise gather. For MLP-like blocks at short context (`SeqLen ≤ 2 × HiddenSize`), GTP_remat-sharding the weight saves *more* memory than recomputing its activations, so the better trade is to keep such modules GTP_remat-sharded and **out** of `--recompute-modules` (offload their activations if needed) — avoiding the third gather entirely. Build the recompute chain only for modules that genuinely need both. + +### 3.2 DDP buckets with (E)GTP_remat + +![DDP + (E)GTP_remat interaction with the distributed optimizer](../../images/generalized_tensor_parallel/0611_ddp_egtp_orthogonal_bucketing.png) + +**(E)GTP_remat is *super loosely coupled* to DDP and the distributed optimizer — they stay completely GTP_remat-agnostic.** GTP_remat is just another sub-axis of the rank grid (`world = TP×GTP_remat×CP×DP`); a GTP_remat-sharded weight rides the *exact same* code path as an ordinary param. There are **no** GTP_remat/EGTP_remat-specific buffers, optimizers, gradient-scaling factors, or bucket groups. The entire DDP/DistOpt stack touches GTP_remat in only **three** narrow places: + +1. **finalize all-reduce** (`_allreduce_replicated_grads_over_gtp_remat_group`) — completes the gtp_remat axis for *replicated* (non-GTP_remat) params (SUM under `calculate_per_token_loss`, AVG otherwise; see §3.2 table); a no-op when GTP_remat is inactive. +2. **`is_gtp_weight_remat` / `allreduce` tags** propagated onto the optimizer's master shards — consumed only by the grad-norm dedup filter. +3. **grad-ready hook routing** (`DistributedDataParallel.__init__`) — for a GTP_remat param, DDP registers its backward post-hook via GTP_remat's `register_grad_accum_hook` instead of autograd's `AccumulateGrad`. GTP_remat fires it from `_handle_megatron_grad_accum` **after** the per-param `{wgrad RS → main_grad add}`. This enforces the invariant below; a no-op (plain autograd path) when GTP_remat is inactive. + +> **Ordering invariant.** A bucket's DDP gradient reduction (the reduce-scatter / all-to-all + local fp32 accumulation) runs **strictly after every GTP_remat param in that bucket has finished `{GTP_remat wgrad RS → main_grad add}`**. `register_grad_ready` only fires the bucket collective once *all* its params are ready, and for GTP_remat params "ready" is signalled by GTP_remat after the add — never by autograd's `AccumulateGrad`, which (because the wgrad RS is async and its `main_grad` accumulation is deferred to a later backward node) can fire **before** the add and would make the bucket reduce read a stale/empty `main_grad` (notably under `reduce_scatter_with_fp32_accumulation`). + +Everything else — bucketing, the reduce-scatter/all-reduce schedule and its overlap, master-state sharding, grad clipping, the checkpoint format — is unchanged and unaware of GTP_remat. + +**Why this matters:** + +- **Free reuse of a mature stack.** GTP_remat inherits DDP's bucketing + comm/compute overlap, the distributed optimizer's fp32-master + Adam-moment sharding, grad-norm/clip, and the existing checkpoint format — no parallel re-implementation to write or maintain (contrast FSDP, which replaces all of these). +- **Orthogonal composability.** Because GTP_remat is a rank-grid sub-axis cut along `out_features` (dim 0, whichever axis TP used), it composes with TP/EP/CP/PP and the DistOpt the same way TP does — no special nesting logic. +- **Zero-cost when off.** With GTP_remat disabled the gtp_remat axis is size-1 and the hooks become no-ops, so non-GTP_remat runs hit byte-identical behavior — GTP_remat can be toggled without forking the DDP/optimizer code paths. +- **Small, auditable surface.** These three hooks are the whole integration contract, which is what makes the correctness argument below tractable. + +DDP groups parameters into **two buffers** by `is_expert_parallel` (MoE tag) — a dense buffer and an expert buffer. GTP_remat/EGTP_remat shards are **merged into** these buffers like ordinary params (no separate GTP_remat/EGTP_remat buckets): they reduce over the replicate group (the default `intra_dp_cp_group` / `intra_expt_dp_group`). + +The DP collective only covers the replicate axis; the gtp_remat axis is completed separately, and **how both axes are scaled depends on the loss normalization** (`config.calculate_per_token_loss`). In all cases each gtp_remat contribution is summed exactly once: + +| | `calculate_per_token_loss=False` (default) | `calculate_per_token_loss=True` | +|---|---|---| +| DDP pre-scale (`gradient_scaling_factor`) | `1/replicate` (= `1/dp_cp_group.size()`) | `1.0` (no pre-scale) | +| gtp_remat reduce-scatter (sharded weights) | **MEAN** (pre-scale wgrad by `1/gtp_remat`) | **SUM** (plain reduce-scatter) | +| finalize over gtp_remat (replicated params) | **AVG** all-reduce | **SUM** all-reduce | +| final normalization | net grad = full `(replicate × gtp_remat)` **mean** | grads summed over all axes, then `÷ total_global_tokens` in `finalize_model_grads` | + +- **Default (mean) path** decouples gradient scaling from the gtp_remat degree: the DP `1/replicate` mean × the reduce-scatter `1/gtp_remat` mean (sharded weights) — or × the finalize AVG (replicated params) — equals the exact full mean, independent of the gtp_remat axis size. +- **Per-token-loss path** must SUM over gtp_remat (like the DP axis): `total_global_tokens` already counts the gtp_remat peers' distinct tokens, so the single `÷ total_global_tokens` does all normalization. A `1/gtp_remat` mean here would shrink every gtp_remat gradient by `1/gtp_remat` (grad-norm mismatch + divergence), so the reduce-scatter mean and finalize AVG are both gated on `not calculate_per_token_loss`. + +> **`average_in_collective` must be off (the default).** The default-path scaling is a *pre-scale* applied before a SUM collective. `average_in_collective=True` instead uses NCCL AVG over the collective's own (replicate) group, which interacts incorrectly with the gtp_remat completion. Asserted via `ProcessGroupCollection.is_gtp_remat_active` in both `arguments.py` (training) and `DistributedDataParallel.__init__` (direct megatron-core users). (Independently, `calculate_per_token_loss` already forbids `average_in_collective`.) + +**Buffer caching.** The per-buffer lists are concatenated once at init into a single flat view for fast iteration in the grad-reduction hot path. + +> **Single distopt instance with GTP_remat.** GTP_remat currently requires `num_distributed_optimizer_instances == 1` (asserted in `parallel_state.py`): partial-distopt sharding of the data domain would need gtp_remat-aware sizing. The dist-opt grad-stats group is therefore the full world. + +### 3.3 Distributed checkpointing (DCP) + +![GTP_remat + DCP save/load reshard for a TP2×GTP2 weight](../../images/generalized_tensor_parallel/0612_gtp_dcp_tp2gtp2_save_load.png) + +GTP_remat supports **PyTorch / Mcore sharded distributed checkpointing** (`--ckpt-format torch_dist`, the `megatron.core.dist_checkpointing` `ShardedTensor` / `ShardedObject` format) for **both model weights and distributed-optimizer state**. Checkpoints are **fully resharding-capable**: a checkpoint saved at one `(TP, GTP_remat, EGTP_remat, DP, PP)` topology can be loaded at a *different* one — including a different GTP_remat/EGTP_remat size — without an offline conversion step. + +Consistent with §3.2, GTP_remat stays *loosely coupled* to the checkpoint stack: there is **no GTP_remat-specific checkpoint format or call path**. The shared `make_sharded_tensors_for_checkpoint` helper became GTP_remat-aware and **delegates internally** to a GTP_remat variant only when the `state_dict` actually contains a `GTPShardedParam` (a no-op otherwise), so call sites are unchanged and non-GTP_remat runs are byte-identical. + +**Save-side call workflow.** The diagram below traces the save path — from `model.sharded_state_dict()` through the `make_*` helpers down to the terminal `ShardedTensor` / `ShardedObject` sinks. The GTP_remat footprint is deliberately tiny: exactly **one new function** (`make_sharded_tensors_for_checkpoint_with_gtp_remat`, in `gtp.py`, which sets `replica_id` for the GTP_remat-*duplicated* entries) plus **one modified function** (the per-tensor `make_tp_sharded_tensor_for_checkpoint` in `core/utils.py`, made GTP_remat-aware in place to emit the GTP_remat-*sharded* offsets). Every other helper is untouched. + +![GTP_remat + DCP checkpoint-save call workflow](../../images/generalized_tensor_parallel/0613_gtp_dcp_save_call_workflow.png) + +**How a GTP_remat weight is described to DCP.** GTP_remat always shards `out_features` (axis 0). The helper layers that GTP_remat split onto the existing TP offsets in the `ShardedTensor`, so the global tensor DCP sees is the *full, unsharded* weight: + +| Weight kind | TP axis | Emitted axis-0 offset | Other axis | +|-------------|---------|------------------------|------------| +| Column-parallel (qkv, fc1) | 0 (same as GTP_remat) | composite `(tp_rank·gtp_remat + gtp_rank, tp·gtp_remat)` | — | +| Row-parallel (proj, fc2) | 1 | GTP_remat-only `(gtp_rank, gtp_remat)` | TP offset on axis 1 | +| No TP (GTP_remat-only) | – | `(gtp_rank, gtp_remat)` | — | + +Because the offsets reconstruct the global shape, the checkpoint is independent of the save-time grid. On load, DCP reads each rank's `[offset : offset+local]` slice from that global and re-tiles it onto the new grid — e.g. `TP1×GTP2`, `TP2×GTP4`, or a DP change. + +**replica_id.** GTP_remat peers hold *distinct* shards (not replicas), so they're disambiguated by their offsets; `replica_id`'s DP coordinate is the GTP_remat-*excluded* replicate rank (one elected writer per shard, per replicate group). **Replicated** tensors that live alongside GTP_remat weights (LayerNorm γ/β, biases, `_extra_state` objects) would otherwise collide across GTP_remat peers, so the helper folds `gtp_rank` into their `replica_id` — exactly one peer is then elected DCP writer per key. + +**`_extra_state`.** This is TransformerEngine's per-module **FP8 calibration state** — for delayed-scaling recipes it holds the `recipe`, the forward/backward `scale` tensors and `amax_history` buffers, plus picklable `extra_fp8_variables`; for BF16 (non-FP8) runs it is an empty tensor. Because it is a pickled byte blob rather than a tensor with a meaningful shape, it is emitted as a `ShardedObject` (via `make_sharded_object_for_checkpoint`), not a `ShardedTensor`. Its amax/scale statistics are *per-tensor globals* for the **full** weight (amax is reduced across the FP8 group), so every GTP_remat peer carries an identical copy — which is exactly why it takes the replicated path above, with `gtp_rank` folded into its `replica_id`. + +**Alignment padding & cross-topology reshard.** When `_gtp_slice_one_param` pads `out_features` to a multiple of `gtp_remat_size · pad_for_alignment`, the saved global describes the *padded* shape, so the helper sets `allow_shape_mismatch=True`. DCP then tolerates a load-side topology whose alignment yields a different padded size — the unpadded data overlaps and the tail pad rows are zeros GTP_remat recomputes. + +> Note: Mamba's `in_proj` is a special case: it **all-gathers its GTP_remat shards** back to the logical TP-local size and strips the pad *before* saving, so its global is topology-independent and needs no `allow_shape_mismatch`. + +**Optimizer state.** The distributed optimizer's master/moment `ShardedObject`s are keyed by `dp_group_idx`. Under GTP_remat/EGTP_remat each peer owns a *different* master shard (the optimizer shards over the gtp_remat/egtp_remat-**excluded** replicate group), so the index is taken from the gtp_remat/egtp_remat-**merged** model-parallel group (`mp_group` for dense, `expt_tp_pp_with_egtp_remat_group` for expert) — giving every peer a distinct key while replicate-group ranks remain true replicas under that key. + +**Pre-save forced param-sync.** Before a save (and around any `disable_forward_pre_hook(param_sync=True)`, e.g. pre-eval), the training loop force-syncs DDP params. `force_param_sync` / `disable_forward_pre_hook` first call `optimizer.prepare_model_params_for_param_sync()`, which copies the FP32 masters into the DDP param buffer, so the sync's `_post_param_sync` copy-back re-quantizes each native-FP8 weight — GTP_remat shards included — from up-to-date masters instead of stale grad scratch under `--reuse-grad-buf-for-mxfp8-param-ag`. The copy-back therefore writes the correct MXFP8 shard, so the forced sync leaves GTP_remat's self-gathered weight intact and does not perturb the next iteration's loss — no GTP-specific preservation is needed. + +### 3.4 Prefetch-chain construction and its design assumptions + +The prefetch chains (§3.1) are **not configured — they are observed at runtime and stored in process-global state**, which imposes assumptions on the weights that every feature combined with GTP_remat must be checked against. + +**Construction (two steps).** + +1. **Classification (once, at build).** `classify_gtp_chains(model)` runs in `training.py`'s `get_model` after the model is built. It walks `named_parameters()` and, for each `GTPShardedParam`, sets `chain_id` (via `_classify_param_chain`, from the active `cuda_graph_modules`) and the dense vs. expert chain. Membership is fixed from here on; re-classifying an already-linked param into a different chain is rejected. + + Routed grouped experts are the exception: their `fc1`/`fc2` weights get their own homogeneous chains for a deeper prefetch — see [Grouped-expert chains](#grouped-expert-chains-one-block-ahead) below. + +2. **Linking (lazily, on the first forward).** The doubly-linked list (`prev_w` / `next_w`) is built the **first time each weight is materialized** inside `all_gather_and_prefetch`: a class-level per-chain cursor (`GTPShardedParam._chain_state[chain_id]["last_weight"]`) records the previously-seen weight, and the current weight links itself after it. The chain therefore **encodes the forward execution order of the first step** and replays it every step after to predict the next weight to prefetch. The recompute chain (`_recompute_next`) self-populates the same way, from the weights re-gathered while `in_fp8_activation_recompute_phase()` is true. + +Weights that must **not** join a chain (embedding, output_layer — they all-gather synchronously and run outside the CUDA-graph boundary) are excluded by setting `weight.prefetch_initialized = True` (and `_need_weight_prefetch = False`) at construction, which skips registration entirely. + +**Why this needs careful consideration.** Because `_chain_state` is a *class attribute* and `prev_w`/`next_w` are strong references between `GTPShardedParam` instances, the chain **holds the weights alive for the life of the process** and **assumes the first step's behavior is representative of every step**. Neither is free: + +| Assumption | What breaks it | Symptom | +|---|---|---| +| **Stable object identity** — `prev_w`/`next_w` point at fixed Python objects | Replacing a weight object at runtime (re-wrapping, checkpoint load that rebinds `.data`, optimizer param swap, resharding) | Chain gathers/prefetches the stale object → wrong weight in the GEMM | +| **Deterministic, fixed forward order** — the observed order is replayed every step | Data-dependent control flow: conditional layers, early exit, MoE routing that skips experts, reordered visitation | Predicted `next_w` is wrong → stale-buffer read or missed prefetch | +| **Single, non-reentrant pass** — one global `last_weight` cursor + per-weight in-flight handles | Two models in one process, an extra autograd graph, unexpected microbatch interleaving | Corrupted cursor / async handles | +| **Fixed, single membership** — `chain_id` and graphed-vs-eager decided once | A weight whose CG scope or dense/expert context changes between steps | Unrepresentable in one linear slot | +| **No parameter sharing/tying** — a linear list gives each weight one slot | A tied/shared param used in two positions (e.g. tied I/O embeddings) | One identity cannot occupy two chain positions; must be excluded | +| **Build-once, run-forever lifetime** — strong refs never released | Building/tearing down GTP models in-process (successive UTs, model re-init, multi-model drivers) | Leaks all GTP params/buffers; a new model's chain can cross-link onto a previous model's stale params | + +**Mitigations.** + +- `reset_gtp_state()` clears the class-level cursors before an in-process rebuild (call it once before `classify_gtp_chains`) — but it does *not* drop `prev_w`/`next_w` links already held by live weights. +- `prefetch_initialized = True` keeps a weight out of the chain — but it is opt-*out* by convention; a new weight that forgets it silently joins. + +**Rule of thumb:** any change that creates/replaces params at runtime, makes forward order data-dependent, runs GTP_remat concurrently, or builds multiple GTP models per process must be checked against the table above. When in doubt, exclude the affected weights so they fall back to synchronous, chain-free all-gather. + +#### Grouped-expert chains (one-block-ahead) + +*Problem.* A chain gives every all-gather exactly **one consume-step of runway** — layer *i*'s AG hides behind layer *i−1*'s GEMM — which suffices only while the transfer stays inside the NVLink domain. Routed-expert weights fail that test twice over: by **volume**, a block gathers `2 × num_experts / EP` expert weights — NCCL-coalesced into just **two** all-gathers, one per role — so those two transfers carry most of the block's bytes; by **distance**, `EGTP_remat` is the group that leaves the NVLink domain. The expert transfer therefore stays partly exposed in **every** MoE block, and the exposure grows as expert count rises and per-GEMM time falls. + +*Design.* When MoE is *not* captured, `linear_fc1` and `linear_fc2` each get their own homogeneous chain (`GTP_remat_grouped_fc1_ungraphed` / `GTP_remat_grouped_fc2_ungraphed`) instead of sharing the general `UNGRAPHED` chain. A homogeneous chain links the **same weight role of consecutive MoE blocks**, so `next_w` points a whole block ahead rather than one GEMM ahead. The roles stay in *separate* chains deliberately: merging them would link `layer_N.fc1 → layer_N.fc2 → layer_{N+1}.fc1 → …`, so `fc1` would prefetch the **same block's** `fc2` — one GEMM of runway again — and only `fc2` would reach across the block boundary. + +*Result.* The win is **resource overlap**, not faster compute and not a faster network: + +- **Runway** — an expert gather now hides behind the entire preceding **MoE block** instead of a single GEMM. +- **Utilization** — the interconnect works under the dense window where it used to idle, and the GPU no longer stalls waiting on the gather: both are busy at once. +- **Cost** — one extra buffer per weight role (see *mandatory double buffering* below). No extra collectives, no change to the math. +- **Bound** — same transfers, same GEMMs, only a different schedule, so the recovered time is exactly the transfer that used to sit on the critical path. + +The figure below puts both schedules on one time axis, aligned at *block start* (**top:** shared chain, **bottom:** per-role chains). **Shaded bands** mark which resource is idle — red where one side waits, green where both are busy; **dashed arrows** trace each gather from the GEMM that launches it to the GEMM that consumes it; the **arrow at the right** is the recovered time, equal to the two hatched `STALL` bars above it. + +![GTP grouped-expert AG prefetch — one-step-ahead vs one-block-ahead](../../images/generalized_tensor_parallel/0725_gtp_grouped_oneblock_prefetch.png) + +Three consequences: +- **One shared stream.** `_stream_key` collapses the fc1/fc2 role, so both chains resolve to a single AG stream and their all-gathers serialize instead of splitting interconnect bandwidth. The capture-axis suffix is preserved, so eager and captured ops still never share a stream. +- **Mandatory double buffering** — this is what makes the deeper prefetch *safe*, and it is not optional: + - the weight cache keys **one buffer per `(shape, dtype, expert_idx)`**, which assumes at most one same-key weight is live; + - one-block-ahead makes block *N* and block *N+1* weights **live at the same time** — same key, two tensors in flight; + - fix: a chain-position **parity (0,1,0,1…)** is folded into the cache key, so consecutive blocks alternate between **exactly two** buffers (counter cleared by `reset_gtp_state()`); + - without it the prefetch would **overwrite the weight the running GEMM is still reading** — a silent-correctness bug, not a crash. +- **Eager only** — the optimization disables itself under CUDA-graph capture: + - `_classify_param_chain` evaluates `graphed = _FULL_ITERATION or ("moe" in cuda_graph_modules)` **before** the split, and returns the plain `GRAPHED` chain when it is true; + - so with `--cuda-graph-impl full_iteration` **every** param is `GRAPHED` — expert weights included — and they keep the ordinary one-step-ahead prefetch; + - why it must: `cuda_graphs.py` drains with `wait_async_comms(GTPChain.GRAPHED.value)`, matching the id **literally**, so a weight in `GTP_remat_grouped_fc1_ungraphed` would never be joined at the graph boundary — a **correctness** hazard, not just a lost overlap; + - lifting it would mean draining by chain-id *prefix* (`_chain_is_grouped`) or registering the grouped streams before capture — neither is done today. + +## 4. Testing + +**Whenever you add or change a GTP_remat/EGTP_remat feature, run the GTP_remat unit-test suite below as a sanity check before opening a PR.** These tests exercise the full TE↔Mcore path (weight gather/RS, DDP, distributed optimizer, finalize, grad-norm) and catch silent-correctness regressions that don't surface as crashes. + +```bash +# 4 GPUs. GTP_remat requires TransformerEngine >= 2.19. +torchrun --nproc-per-node 4 -m pytest tests/unit_tests/generalized_tensor_parallel/ -v +``` + +| Test file | What it guards | +|-----------|----------------| +| `test_gtp_basics.py` | Core GTP_remat shard/gather + DDP bucket alignment. | +| `test_attention_gtp.py` | GTP_remat on attention linears, loss parity vs no-GTP_remat. | +| `test_mamba_gtp.py` | GTP_remat on Mamba projection weights. | +| `test_tp_gtp.py` | GTP_remat composed with tensor parallelism (`tp_group × gtp_remat_group`). | +| `test_moe_egtp.py` | EGTP_remat on MoE routed-expert weights. | +| `test_gtp_loss_correctness.py` | End-to-end: GTP_remat per-step loss trajectory matches a no-GTP_remat baseline. | +| `test_gtp_grad_correctness.py` | Gradient + dist-opt + grad-norm numeric parity vs a DP baseline at replicate (DP) > 1. | +| `test_gtp_cudagraph_grad.py` | Capture-step grad-norm guard (§1.2): `_backup_grads_before_capture`/`_restore_grads_after_capture` keep a graph capture from clobbering finalized `main_grad` (own params + cross-graph `next_w`, incl. routed-expert `weight_list`). | +| `test_gtp_dcp.py` | DCP sharding metadata (§3.3): TP×GTP_remat offsets, pad reshard, `replica_id`, native-FP8 save/load. | +| `test_gtp_muon_dcp.py` | Muon optimizer-state DCP roundtrip (§1.6): `replica_id` fold + native-FP8 backfill matching. | +| `test_gtp_fp8_param_gather.py` | Native-FP8 GTP_remat (§1.3): fp8-vs-BF16 loss parity (TP1/TP2, MoE), post-save-spike guard. | + +All tests require ≥ 4 GPUs and TransformerEngine >= 2.19; they self-skip when those are unavailable. A green run (skips for unmet hardware/config are acceptable) is the minimum bar for any GTP_remat change. diff --git a/docs/images/generalized_tensor_parallel/0611_ddp_egtp_orthogonal_bucketing.png b/docs/images/generalized_tensor_parallel/0611_ddp_egtp_orthogonal_bucketing.png new file mode 100644 index 00000000000..2d311138e8d Binary files /dev/null and b/docs/images/generalized_tensor_parallel/0611_ddp_egtp_orthogonal_bucketing.png differ diff --git a/docs/images/generalized_tensor_parallel/0612_gtp_dcp_tp2gtp2_save_load.png b/docs/images/generalized_tensor_parallel/0612_gtp_dcp_tp2gtp2_save_load.png new file mode 100644 index 00000000000..937846e9f0b Binary files /dev/null and b/docs/images/generalized_tensor_parallel/0612_gtp_dcp_tp2gtp2_save_load.png differ diff --git a/docs/images/generalized_tensor_parallel/0613_gtp_dcp_save_call_workflow.png b/docs/images/generalized_tensor_parallel/0613_gtp_dcp_save_call_workflow.png new file mode 100644 index 00000000000..b69bd835769 Binary files /dev/null and b/docs/images/generalized_tensor_parallel/0613_gtp_dcp_save_call_workflow.png differ diff --git a/docs/images/generalized_tensor_parallel/0617_gtp64_weak_scaling_efficiency.png b/docs/images/generalized_tensor_parallel/0617_gtp64_weak_scaling_efficiency.png new file mode 100644 index 00000000000..03fc587f96a Binary files /dev/null and b/docs/images/generalized_tensor_parallel/0617_gtp64_weak_scaling_efficiency.png differ diff --git a/docs/images/generalized_tensor_parallel/0628_gtp_remat_class_hierarchy.png b/docs/images/generalized_tensor_parallel/0628_gtp_remat_class_hierarchy.png new file mode 100644 index 00000000000..e98dca8d480 Binary files /dev/null and b/docs/images/generalized_tensor_parallel/0628_gtp_remat_class_hierarchy.png differ diff --git a/docs/images/generalized_tensor_parallel/0712_gtp_te_protocol_redesign.png b/docs/images/generalized_tensor_parallel/0712_gtp_te_protocol_redesign.png new file mode 100644 index 00000000000..479e39b3f57 Binary files /dev/null and b/docs/images/generalized_tensor_parallel/0712_gtp_te_protocol_redesign.png differ diff --git a/docs/images/generalized_tensor_parallel/0725_gtp_grouped_oneblock_prefetch.png b/docs/images/generalized_tensor_parallel/0725_gtp_grouped_oneblock_prefetch.png new file mode 100644 index 00000000000..5cacd7f20f9 Binary files /dev/null and b/docs/images/generalized_tensor_parallel/0725_gtp_grouped_oneblock_prefetch.png differ diff --git a/docs/index.md b/docs/index.md index 11337315588..3995c70217d 100644 --- a/docs/index.md +++ b/docs/index.md @@ -66,11 +66,9 @@ models/index :caption: Advanced Features user-guide/features/moe -user-guide/features/context_parallel user-guide/features/megatron_fsdp user-guide/features/dist_optimizer user-guide/features/optimizer_cpu_offload -user-guide/features/pipeline_parallel_layout user-guide/features/fine_grained_activation_offloading user-guide/data-loading user-guide/features/megatron_energon @@ -103,5 +101,6 @@ apidocs/index.rst :hidden: :caption: Resources +user-guide/hybrid-model-migration advanced/index ``` diff --git a/docs/user-guide/features/context_parallel.md b/docs/user-guide/features/context_parallel.md index 890609ac7de..31965a2fbdb 100644 --- a/docs/user-guide/features/context_parallel.md +++ b/docs/user-guide/features/context_parallel.md @@ -7,9 +7,7 @@ license agreement from NVIDIA CORPORATION is strictly prohibited. --> -# Context Parallel Package - -## Context Parallelism Overview +# Context Parallel Overview ```{figure} ../../images/context_parallel/CP_overview.png :alt: Diagram of a transformer layer with tensor parallelism 2 and context parallelism 2, showing CP and TP communication patterns around attention and other blocks. @@ -40,4 +38,3 @@ CP addresses these tradeoffs. With CP, each GPU computes on part of the sequence CP support is included on the GPT code path. Other models that share that path, such as LLaMA, can use CP as well. CP works with TP (tensor model parallelism), PP (pipeline model parallelism), and DP (data parallelism). The total GPU count is TP × CP × PP × DP. CP also works with different attention variants, including MHA, MQA, and GQA, with unidirectional or bidirectional masking. Enable CP by setting `context_parallel_size=` on the command line. The default `context_parallel_size` is 1, which disables CP. Running with CP requires Megatron Core (>=0.5.0) and Transformer Engine (>=1.1). - diff --git a/docs/user-guide/features/index.md b/docs/user-guide/features/index.md index cb2e895afdc..ff805d67eb2 100644 --- a/docs/user-guide/features/index.md +++ b/docs/user-guide/features/index.md @@ -17,12 +17,10 @@ Guides for Megatron Core training features. cuda_graph fine_grained_activation_offloading moe -context_parallel megatron_fsdp dist_optimizer optimizer_cpu_offload paged_stash -pipeline_parallel_layout tokenizers megatron_energon megatron_rl diff --git a/docs/user-guide/hybrid-model-migration.md b/docs/user-guide/hybrid-model-migration.md new file mode 100644 index 00000000000..8f9d02bf18f --- /dev/null +++ b/docs/user-guide/hybrid-model-migration.md @@ -0,0 +1,400 @@ + + +# Migrate from GPTModel to HybridModel + +This guide describes how to replace a Megatron Core `GPTModel` with a +`HybridModel`, convert an existing distributed checkpoint, and start or resume +training with the converted weights. The conversion stays in Megatron's +distributed-checkpoint format; it does not use Hugging Face as an intermediate +format. + +## 1. What Is HybridModel? + +A standard `GPTModel` decoder layer contains both a self-attention sublayer and +an MLP or MoE sublayer under one layer index. `HybridModel` instead builds an +ordered stack in which every position represents one layer family. The order is +described by `--hybrid-layer-pattern`: + +| Symbol | Layer family | +|--------|--------------| +| `M` | Mamba-2 state-space layer | +| `G` | Gated Delta Network (GDN) layer | +| `*` | Self-attention layer | +| `D` | DeepSeek Sparse Attention (DSA) layer | +| `-` | Dense MLP layer | +| `E` | Mixture-of-Experts (MoE) layer | + +One pattern symbol is one HybridModel layer. Consequently, one GPT transformer +block becomes two HybridModel layers when preserving the original architecture: + +| Source architecture | Equivalent HybridModel pattern | Hybrid layer count | +|---------------------|--------------------------------|--------------------| +| Two dense GPT blocks | `*-*-` | 4 | +| Two all-layer MoE GPT blocks | `*E*E` | 4 | + +For example, source GPT layer 0 is split between Hybrid layers 0 and 1: +its attention parameters move to the first `*`, and its MLP parameters move to +the first `-` or `E`. Source GPT layer 1 maps to the next pair, and so on. + +The pattern can also describe execution layout. A `|` marks a pipeline segment +boundary, and `/` introduces a repeated Multi-Token Prediction (MTP) pattern. +For example, `*-*-|*-*-` places four GPT-equivalent blocks across two pipeline +segments. Separators do not count as layers. + +HybridModel provides the following benefits: + +- Different layer families can be composed in one model without forcing every + decoder block to have the same structure. +- Attention and dense or expert MLP layers can be placed and configured + independently. +- Mamba, GDN, standard or DeepSeek attention, dense MLP, and MoE layers can use + one pattern-driven model interface. +- Patterns that include Mamba can replace some quadratic attention layers with + subquadratic sequence mixing and fixed-size recurrent inference state. +- Pipeline and virtual-pipeline segmentation can be expressed with the model + pattern instead of a separate layer layout. +- The same model abstraction can describe a pure transformer (`*-` repeated), + a pure Mamba model, or a heterogeneous architecture. + +These capabilities do not imply an automatic throughput or quality improvement. +An architecture-preserving `*-` or `*E` migration should be validated for +numerical equivalence, and a pattern that adds another layer family should be +treated as a new architecture and benchmarked independently. + +## 2. How to Convert a Checkpoint + +There are two ways to bring `GPTModel` weights into a `HybridModel` run. Both +stay in Megatron's distributed-checkpoint format and can reshard across a +different tensor, pipeline, expert, or FSDP layout on the following load. + +- **Option A — translate at load time (no separate step).** Start the hybrid + run directly against the GPT checkpoint. The hybrid model retargets its own + checkpoint state dict at the GPT checkpoint's keys during loading, so no + second copy is written to disk. This supports both `torch_dist` checkpoints + and Megatron-FSDP `fsdp_dtensor` checkpoints, including their optimizer + state. This path also supports patterns that contain layer families with no + GPT counterpart, such as Mamba (`M`) positions, which keep their fresh + initialization. +- **Option B — convert offline to a new checkpoint.** Use + [`tools/checkpoint/gpt_hybrid_conversion.py`](https://github.com/NVIDIA/Megatron-LM/blob/main/tools/checkpoint/gpt_hybrid_conversion.py) + to write a standalone `HybridModel` checkpoint whose keys already match the + hybrid layout. Use this when you need a persisted hybrid checkpoint, an + architecture-preserving `*-` or `*E` copy, or a target you can inspect before + training. This path only supports `*-` and `*E` layouts. + +### Option A: Translate at load time + +Load-time translation is handled by +[`megatron/core/dist_checkpointing/gpt_checkpoint_interop.py`](https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/core/dist_checkpointing/gpt_checkpoint_interop.py). +It triggers automatically when a non-hybrid (GPT) checkpoint is loaded into a +`HybridModel` run: for `torch_dist`, the run's model and optimizer sharded +state dicts are rewritten into the GPT checkpoint's homogeneous-layer format; +for `fsdp_dtensor`, their explicit parameter-name mappings are rewritten onto +the GPT keys before Torch DCP planning. The checkpoint is read directly, and +the weights and optimizer state are resharded to the current +TP/PP/EP/ETP/FSDP layout. No conversion tool is run, and the GPT checkpoint on +disk is never modified. + +The reverse mismatch is an error: loading a checkpoint that was saved by a +hybrid run into a non-hybrid run raises a `RuntimeError` that directs you to the +hybrid training entrypoint. + +#### Select the checkpoint semantics + +When the hybrid run loads a GPT checkpoint, it must set +`--hybrid-layer-pattern` so checkpoint layers can be paired with hybrid layer +positions. Point either `--load` or `--pretrained-checkpoint` at the GPT +checkpoint root (see [Section 3](#3-how-to-train-a-model)). + +`--finetune` is optional and retains its normal checkpoint-loading meaning; the +GPT-to-Hybrid translation does not select it on the user's behalf: + +- Without `--finetune`, a direct `--load` resumes iteration, optimizer, + scheduler, RNG, and rerun state according to the normal checkpoint and + parallel-layout compatibility rules. +- With `--finetune`, iteration, scheduler, RNG, and rerun state restart fresh. + Model weights still load, and translated optimizer state loads unless + `--no-load-optim` is set. +- The existing `--pretrained-checkpoint` fallback uses finetuning semantics + when the `--load` directory contains no checkpoint. Use `--load` directly + when full resume semantics are desired. + +By default the GPT run's **optimizer state is also translated and loaded** — +Adam moments and fp32 master params for the attention and MLP layers carry over, +enabling architecture-preserving continued training. Pass `--no-load-optim` to +skip this and start every layer's optimizer state fresh. + +For `torch_dist`, loading optimizer state requires the GPT checkpoint to use a +model-space distributed-optimizer format (`fully_reshardable` or +`fully_sharded_model_space`, i.e. saved with +`--dist-ckpt-optim-fully-reshardable`). The bucket-space formats key optimizer +state by a flat buffer layout that the extra hybrid layers reshuffle, so the +run raises an error directing you to re-save the checkpoint or pass +`--no-load-optim`. + +For Megatron FSDP, save and load with `--ckpt-format fsdp_dtensor`. The loader +retargets the explicit DTensor model keys and the model-parameter names used by +the distributed optimizer, so model weights and optimizer state can be +resharded across a different FSDP, TP, EP, or ETP layout during the automatic +GPT-to-Hybrid load. This path is for Megatron FSDP; Torch FSDP2's `torch_dcp` +format is not supported by this automatic translation. + +Layers without a GPT counterpart (for example Mamba `M` positions) have no +optimizer state in the checkpoint; their moments start fresh, and the run prints +a warning naming how many layers are affected. + +#### Supported patterns and key mapping + +The main pattern (the part before any `/` MTP suffix, with `|` pipeline +separators ignored) may contain: + +| Symbol | Source of weights | +|--------|-------------------| +| `*` | GPT `self_attention` sub-module of the paired layer | +| `-` or `E` | GPT `mlp` sub-module of the paired layer (MoE tensors also live under `mlp.*`) | +| `M` | No GPT source; the Mamba layer keeps its fresh initialization | + +Parameters are paired by occurrence: the *i*-th `*` position takes GPT layer +*i*'s attention, and the *i*-th `-`/`E` position takes GPT layer *i*'s MLP. +`decoder.final_norm` is loaded from GPT's `decoder.final_layernorm`, and +embedding and output weights are copied unchanged. + +Because each GPT layer supplies exactly one attention and one MLP sub-module, +the loader rejects a pattern that: + +- contains MTP layers (a `/...` suffix), which have no GPT source weights; +- uses a layer type it cannot translate, such as GDN (`G`) or DeepSeek Sparse + Attention (`D`), whose weight layouts differ from GPT attention; +- mixes dense (`-`) and MoE (`E`) MLP positions in one pattern; or +- has an unequal or zero number of `*` and MLP positions. + +The checkpoint's `num_layers` must equal the number of `*` positions in the +pattern; a mismatch is rejected. + +```{warning} +Optimizer translation covers the Adam moments and fp32 master params. Pass +`--no-load-optim` for a weights-only load. Use `--finetune` when iteration, +scheduler, RNG, and rerun state should restart instead of resume. +``` + +### Option B: Convert offline with `gpt_hybrid_conversion.py` + +#### Choose an architecture-preserving pattern + +For a source checkpoint with *N* GPT layers: + +- Use `*-` repeated *N* times for a dense GPT model. +- Use `*E` repeated *N* times for a GPT model whose MLP in every layer is MoE. + +The converter maps parameters by occurrence, not merely by numeric layer index: + +| Source parameter | Target parameter | +|------------------|------------------| +| Attention from GPT layer *i* | The *i*-th `*` layer | +| MLP or MoE from GPT layer *i* | The *i*-th `-` or `E` layer | +| Embedding and output weights | Copied without changing their model role | +| `decoder.final_layernorm` | Renamed to `decoder.final_norm` | + +#### Check the prerequisites + +The source must use one of these distributed-checkpoint formats: + +- `torch_dist` +- `fsdp_dtensor` + +Prefer a top-level checkpoint root containing +`latest_checkpointed_iteration.txt` as `--load-dir`. If `--load-dir` points +directly to a directory containing `metadata.json`, the converter writes a flat +target without a tracker file; the standard training entry point expects a +checkpoint root and tracker. + +Run the converter from the repository root in a Megatron environment. A plain +`python` process is sufficient; `torchrun` and a GPU are not required. The tool +gathers full logical tensors on CPU, so the host must have enough memory for the +unsharded source and target model state dicts. Always use a target directory +that is different from the source directory. + +```{warning} +This is a weights conversion, not a resumable full training-state conversion. +Sharded optimizer, RNG, rerun, and Transformer Engine `_extra_state` tensors are +not converted. Some non-tensor entries can remain in `common.pt`, but they do +not constitute a converted optimizer or RNG state. Start the converted model +with a fresh optimizer and RNG state. +``` + +#### Run the conversion + +The following example converts a four-layer dense GPT model. Its equivalent +HybridModel has the eight-layer pattern `*-*-*-*-`: + +```bash +uv run python tools/checkpoint/gpt_hybrid_conversion.py \ + --direction gpt-to-hybrid \ + --load-dir /path/to/gpt-checkpoints \ + --save-dir /path/to/hybrid-checkpoints \ + --hybrid-layer-pattern '*-*-*-*-' \ + --reset-iterations +``` + +Always quote the pattern because `*` and `|` have special meaning to a shell. +`--input-format auto` and `--output-format auto` are the defaults: the tool +detects the source backend and writes the same backend. `--reset-iterations` +resets the checkpoint iteration, consumed-sample counters, and cached +`train_iters` and `train_samples`; omit it when the new run must retain that +schedule metadata. + +The number of `*` positions and the number of `-` or `E` positions must both +equal the source GPT layer count. The pattern validator rejects GDN, DSA, and +mixed dense/MoE layouts. When cached training arguments are present, the tool +also rejects interleaved MoE, experimental or linear attention, heterogeneous +block specifications, Multi-Latent Attention, and MTP checkpoints. That +source-feature validation is incomplete when `common.pt` has no cached `args` +or an older checkpoint lacks a field, so verify those features manually. + +The conversion recognizes standard attention and MLP/MoE state-dict keys only. +Other layer-local tensors are omitted. The documented `hybrid_stack_spec` also +uses Transformer Engine's fused layernorm/linear layout; a local or otherwise +non-TE source layout requires a compatible custom Hybrid stack and key +conversion. Always perform the strict-load check described below. + +Do not append an MTP `/...` suffix during conversion. The converter only maps +the main pattern before the first `/`, so it does not create MTP parameters. + +When the source path is a checkpoint root with +`latest_checkpointed_iteration.txt`, the output contains an iteration directory +and a matching tracker file. The saved full-shape tensors can be resharded by a +later Megatron load for a different tensor, pipeline, expert, or FSDP layout. + +## 3. How to Train a Model + +### Update the training command + +Start with the command that trained the GPT model and make these changes: + +1. Replace `pretrain_gpt.py` with `pretrain_hybrid.py`. +2. Remove `--num-layers` and add the same ordered main-layer symbols used for + conversion. Pipeline `|` separators may be added or moved. The command-line + parser derives `num_layers` from the pattern. +3. Select the HybridModel stack specification with + `--spec megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec`. +4. Point a checkpoint input at the pretrained weights and write new training + checkpoints to a separate directory: + - With **Option A (load-time translation)**, point + `--load` directly at the *GPT* checkpoint for resume semantics, or use + `--pretrained-checkpoint` for finetuning semantics. The optimizer state is + loaded by default; add `--no-load-optim` only if you want a fresh + optimizer. No offline conversion is needed, and `--finetune` is not + required by the translation. + - With **Option B (offline conversion)**, point `--pretrained-checkpoint` at + the converted *hybrid* checkpoint. Set `--ckpt-format` to the converter's + `torch_dist` or `fsdp_dtensor` output format. + +A minimal Option A migration — loading the GPT checkpoint and its optimizer +state directly for architecture-preserving continued training — looks like this: + +```diff +- torchrun --nproc_per_node=8 pretrain_gpt.py \ +- --num-layers 4 \ +- --load /path/to/gpt-checkpoints \ +- --save /path/to/gpt-checkpoints ++ torchrun --nproc_per_node=8 pretrain_hybrid.py \ ++ --hybrid-layer-pattern '*-*-*-*-' \ ++ --spec megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec \ ++ --load /path/to/gpt-checkpoints \ # first launch; switch to the save directory later ++ --save /path/to/new-training-checkpoints +``` + +For Option B, point `--pretrained-checkpoint` at the converted hybrid +checkpoint instead. + +Keep the existing architecture, optimizer, precision, data, and basic +TP/DP/EP/CP arguments unless this guide identifies a required change. Review +pattern-driven pipeline layout and GPT-specific dataset features separately. + +```{warning} +`pretrain_hybrid.py` does not select `GPTFIMDataset` when `--fim-data` is set. +A GPT training workflow that uses fill-in-the-middle data needs a custom dataset +path or equivalent Hybrid entry-point support before migration. +``` + +With an empty `--load` directory, `--pretrained-checkpoint` loads the pretrained +weights with finetuning semantics: iteration starts at zero and RNG state is not +restored. For Option A the optimizer state is still warm-started unless +`--no-load-optim` is set (Option B always starts with a fresh optimizer). After +the job writes a checkpoint to `--load`, later launches resume the new +HybridModel training state normally. + +### Train from scratch + +To initialize every layer from scratch, use the same `pretrain_hybrid.py`, +`--hybrid-layer-pattern`, and `--spec` arguments, but omit +`--pretrained-checkpoint`. Point `--load` and `--save` at the new run directory +if later launches should resume it. Unlike checkpoint conversion, training from +scratch can use compatible layer families supported by the selected HybridModel +stack specification and can include an MTP suffix such as `M*M*/MM/MM`. +Pattern constraints still apply; for example, standard attention `*` and DSA +`D` cannot be used in the same model. + +### Account for expanded layer indices + +Any setting, mapping, or callback indexed by decoder layer must use HybridModel +indices. For the pattern `*E*E`, source layer 0 attention is Hybrid layer 0 and +its MoE is Hybrid layer 1; source layer 1 attention is Hybrid layer 2 and its +MoE is Hybrid layer 3. Expand attention-only lists with inactive entries for +the intervening MLP or MoE positions. + +### Configure pipeline parallelism + +For pipeline parallelism, add `|` separators without changing the ordered layer +symbols. For example, the converted pattern `*-*-*-*-` can be trained with two +pipeline segments as `*-*-|*-*-`. The number of pipe-delimited segments must be +divisible by `--pipeline-model-parallel-size`. + +The pattern replaces conventional pipeline layout controls. Remove +`--num-layers-per-virtual-pipeline-stage`, +`--num-virtual-stages-per-pipeline-rank`, `--pipeline-model-parallel-layout`, +`--account-for-embedding-in-pipeline-split`, and +`--account-for-loss-in-pipeline-split`. When the pattern contains `|`, also +remove `--decoder-first-pipeline-num-layers` and +`--decoder-last-pipeline-num-layers`. Express virtual-pipeline segmentation +with additional pipe-delimited segments instead. + +The declarative `HybridModelBuilder` currently rejects virtual pipeline +parallelism. Pipe-defined virtual stages are supported by the +`pretrain_hybrid.py` CLI builder, but custom builder users must avoid VPP or use +a path that explicitly supports it. + +### Update custom providers and conversion mappings + +Custom providers and conversion mappings also need to account for these API and +state-dict differences: + +- Build or register `HybridModel` instead of `GPTModel`. +- Supply a `hybrid_stack_spec` instead of a GPT transformer-layer spec. +- Set programmatic `num_layers` to the number of layer symbols in the main + pattern; unlike the CLI path, a custom provider might not derive it. +- Map attention and MLP/MoE parameters to their separate Hybrid layer indices. +- Use `decoder.final_norm` in HybridModel mappings instead of + `decoder.final_layernorm`. +- Expand per-layer settings such as attention-window schedules to the full + HybridModel pattern. + +### Validate before scaling up + +Before starting a long run: + +- Load the converted checkpoint strictly and confirm that no model keys or + tensor shapes are missing or unexpected. +- For a `*-` or `*E` migration, compare logits on a fixed batch against the + source GPT model within the expected precision tolerance. +- Run a few training iterations and inspect the loss, gradient norms, and + parameter counts by layer. +- Save and reload one new checkpoint to confirm that the new optimizer and RNG + state resume correctly. diff --git a/docs/user-guide/parallelism-guide.md b/docs/user-guide/parallelism-guide.md index 2540ca0a827..375afde60f4 100644 --- a/docs/user-guide/parallelism-guide.md +++ b/docs/user-guide/parallelism-guide.md @@ -11,6 +11,14 @@ Megatron Core supports multiple parallelism strategies that can be combined to efficiently train models from billions to trillions of parameters across thousands of GPUs. +```{toctree} +:hidden: + +features/context_parallel +features/pipeline_parallel_layout +../api-guide/core/generalized_tensor_parallel +``` + ## Overview The following table summarizes supported parallelism strategies. @@ -116,7 +124,7 @@ Split long sequences across GPUs for efficient long-context training. - Reduces activation memory - Can combine with TP, PP, DP -Refer to [Context Parallelism Deep Dive](features/context_parallel.md) for a detailed guide with performance analysis. +Refer to the [Context Parallel Overview](features/context_parallel.md) for a detailed guide with performance analysis. ## Expert Parallelism (EP) diff --git a/examples/inference/advanced/gpt_dynamic_inference.py b/examples/inference/advanced/gpt_dynamic_inference.py index 21cae1792e4..269538854d6 100644 --- a/examples/inference/advanced/gpt_dynamic_inference.py +++ b/examples/inference/advanced/gpt_dynamic_inference.py @@ -120,7 +120,7 @@ def _add_request(): nonlocal num_requests_added _request = requests[num_requests_added] engine.add_request(num_requests_added, _request.prompt_text, _request.sampling_params) - _request.time_start = get_curr_time() + _request.time_start = get_curr_time(do_broadcast=False) _request.state = "started" num_requests_added += 1 tbar.update(1) @@ -129,7 +129,10 @@ def _process_step_result(result): """Process a single engine step result, updating bookkeeping state.""" nonlocal total_output_tokens, num_requests_finished - is_decode_only = engine.is_decode_only + decode_only = engine.decode_only + is_decode_only = ( + decode_only.launched if decode_only.launched is not None else decode_only.consumed + ) # Record cuda_graph_request_count. cuda_graph_request_count = result["cuda_graph_request_count"] @@ -149,14 +152,14 @@ def _process_step_result(result): step_times["prefill"].append(step_time) # Append output tokens. - output_start = get_curr_time() + output_start = get_curr_time(do_broadcast=False) for finished_request_record in finished_request_records: finished_request = finished_request_record.merge() # Update local request object. request = requests[finished_request.request_id] - request.time_end = get_curr_time() + request.time_end = get_curr_time(do_broadcast=False) request.state = "finished" request.request_id = finished_request.request_id request.events = finished_request.events @@ -186,16 +189,16 @@ def _process_step_result(result): if not finished_request.sampling_params.skip_prompt_log_probs: request.prompt_top_n_logprobs = finished_request.prompt_top_n_logprobs num_requests_finished += 1 - output_times.append(get_curr_time() - output_start) + output_times.append(get_curr_time(do_broadcast=False) - output_start) if batch_ranges is not None: # Batch-drain mode: add all requests in a batch, drain, then next batch. for batch_idx, (batch_start, batch_end) in enumerate(batch_ranges): # Add all requests in current batch. - add_start = get_curr_time() + add_start = get_curr_time(do_broadcast=False) while num_requests_added < batch_end: _add_request() - add_times.append(get_curr_time() - add_start) + add_times.append(get_curr_time(do_broadcast=False) - add_start) # Step until all active requests finish (drain). while engine.has_unfinished_requests(): @@ -213,7 +216,7 @@ def _process_step_result(result): # Original mode: add requests per step based on arrival time or count. while True: # Add requests. - add_start = get_curr_time() + add_start = get_curr_time(do_broadcast=False) if args.incoming_requests_per_step is None: # Add requests with 'earlier' arrival time. while num_requests_added < num_requests_total: @@ -226,10 +229,10 @@ def _process_step_result(result): min(args.incoming_requests_per_step, num_requests_total - num_requests_added) ): _add_request() - add_times.append(get_curr_time() - add_start) + add_times.append(get_curr_time(do_broadcast=False) - add_start) # Step inference engine (i.e., generate a token for each active request). - # Before step, we haven't done the scheduling, so we cannot know the is_decode_only + # The engine reports the consumed and launched decode-only states after scheduling. try: result = engine.step_modern() except EngineSuspendedError as e: @@ -476,7 +479,9 @@ def escape_str(s): # Attach peak memory metrics; the functional test only validates these # if the fields exist in the golden values. json_results.update(peak_mem_stats) - json_results["lifetime_prefill_token_count"] = engine.context.lifetime_prefill_token_count + json_results["lifetime_prefill_token_count"] = ( + engine.context.lifetime_prefill_token_count + ) json_results["async_sched_step_count"] = engine.context.async_sched_step_count json_results["async_sched_compaction_step_count"] = ( engine.context.async_sched_compaction_step_count diff --git a/examples/inference/utils.py b/examples/inference/utils.py index 104c1d4b201..cb1a5dd11f9 100644 --- a/examples/inference/utils.py +++ b/examples/inference/utils.py @@ -34,10 +34,24 @@ def get_default_sampling_params(termination_id: int = None): def get_curr_time(do_broadcast: bool = True) -> float: - """Get synchronized time across ranks.""" + """Get the current time, optionally synchronized across distributed ranks. + + Args: + do_broadcast (bool): Whether multi-rank callers require a rank-zero + timestamp broadcast. + + Returns: + float: Current time in seconds. + """ + if ( + not do_broadcast + or not torch.distributed.is_initialized() + or torch.distributed.get_world_size() == 1 + ): + return time.time_ns() / 10**9 + curr_time = torch.cuda.LongTensor([time.time_ns()]) - if torch.distributed.is_initialized() and do_broadcast: - torch.distributed.broadcast(curr_time, src=0) + torch.distributed.broadcast(curr_time, src=0) return curr_time.item() / 10**9 @@ -345,7 +359,10 @@ def print_unique_prompts_and_outputs(results: List["DynamicInferenceRequest"]) - unique_prompt_map[req.prompt].append(idx) for unique_idx, (prompt_text, request_idxs) in enumerate(unique_prompt_map.items()): - prompt_len = len(results[request_idxs[0]].prompt_tokens) + request = results[request_idxs[0]] + prompt_len = request.prompt_length + if prompt_len is None and request.prompt_tokens is not None: + prompt_len = len(request.prompt_tokens) print( f"\n{unique_idx+1}/{len(unique_prompt_map)}" f"[n {len(request_idxs)}, l {prompt_len}] {escape_str(prompt_text)}" @@ -401,7 +418,9 @@ def dump_inference_results_to_json( lifetime_prefill_token_count (int): Total prefill tokens processed. async_sched_step_count (int): Number of async scheduling decode steps. async_sched_compaction_step_count (int): Number of async scheduling decode - steps where post-forward compaction discarded finished rows. + steps that discarded speculative rows for finished requests. This + includes identity-prefix and all-finished cases that require no GPU + gather. """ if not args.output_path: return @@ -442,9 +461,7 @@ def dump_inference_results_to_json( json_results.update(peak_mem_stats) json_results["lifetime_prefill_token_count"] = lifetime_prefill_token_count json_results["async_sched_step_count"] = async_sched_step_count - json_results["async_sched_compaction_step_count"] = ( - async_sched_compaction_step_count - ) + json_results["async_sched_compaction_step_count"] = async_sched_compaction_step_count print(f' Saving results to {args.output_path}') with open(args.output_path, "w") as fp: diff --git a/examples/mimo/pretrain_mimo.py b/examples/mimo/pretrain_mimo.py index 4f304958f88..ee521b188a5 100644 --- a/examples/mimo/pretrain_mimo.py +++ b/examples/mimo/pretrain_mimo.py @@ -5,6 +5,7 @@ from __future__ import annotations import argparse +from functools import partial from examples.mimo.model_providers import resolve_provider from examples.mimo.model_providers.nemotron_moe_vlm import add_model_provider_args @@ -16,9 +17,16 @@ from examples.mimo.training.builder import MimoBuildConfig from examples.mimo.training.data import add_mock_data_args, build_train_valid_test_data_loaders from examples.mimo.training.distributed import initialize_distributed, shutdown_distributed +from examples.mimo.training.encoder_prefetch import ( + EncoderPrefetchLoader, + add_encoder_prefetch_args, + prefetch_frozen_features, + validate_encoder_prefetch_args, +) from examples.mimo.training.step import mimo_forward_step from examples.mimo.training.topology import create_topology from megatron.core.enums import ModelType +from megatron.core.utils import unwrap_model from megatron.training.argument_utils import pretrain_cfg_container_from_args from megatron.training.arguments import parse_args, validate_args from megatron.training.global_vars import set_global_variables @@ -31,6 +39,7 @@ def extra_args_provider(parser: argparse.ArgumentParser) -> argparse.ArgumentPar parser = add_model_provider_args(parser) parser = add_hetero_grid_args(parser) parser = add_mock_data_args(parser) + parser = add_encoder_prefetch_args(parser) return parser @@ -54,6 +63,7 @@ def _parse_and_validate() -> argparse.Namespace: args.world_size = physical_world_size if not args.use_distributed_optimizer: raise ValueError("heterogeneous MIMO training requires --use-distributed-optimizer") + validate_encoder_prefetch_args(args) if getattr(args, "padded_vocab_size", None) is None: args.padded_vocab_size = calculate_padded_vocab_size( @@ -68,43 +78,74 @@ def main() -> None: set_global_variables(args, build_tokenizer=False) provider = resolve_provider(args) - topology = None - try: - initialize_distributed() - # The grid/rank-layout args model a single encoder region; the builder itself is - # generic over any number of encoder grids in the topology. - encoder_name = provider.encoder_module_names[0] if provider.encoder_module_names else None - specs = build_module_grid_specs(args, args.world_size, encoder_name) - topology = create_topology(specs) + prefetch_loader = None + initialize_distributed() + # The grid/rank-layout args model a single encoder region; the builder itself is + # generic over any number of encoder grids in the topology. + encoder_name = provider.encoder_module_names[0] if provider.encoder_module_names else None + specs = build_module_grid_specs(args, args.world_size, encoder_name) + topology = create_topology(specs) - communicator = provider.build_communicator(args, topology) + communicator = provider.build_communicator(args, topology) - loaders = build_train_valid_test_data_loaders(args, topology) - iterators = tuple(iter(loader) if loader is not None else None for loader in loaders) + if args.mimo_encoder_prefetch and len(provider.encoder_module_names) != 1: + raise ValueError("encoder prefetch requires exactly one encoder") + + # Encoder prefetch runs encoder forward while producing batches, so it needs the built + # rank-local encoder instance. Capture the wrapped model here so the data provider can + # later extract that encoder and bind it to the prefetch worker. + captured_model = {} + hooks = [] + if args.mimo_encoder_prefetch: + + def capture_model(models): + captured_model["model"] = models[0] + return models - model_cfg = MimoBuildConfig(_topology=topology) - cfg = pretrain_cfg_container_from_args(args, model_cfg) + hooks.append(capture_model) + model_cfg = MimoBuildConfig(_topology=topology, post_wrap_hooks=hooks) + cfg = pretrain_cfg_container_from_args(args, model_cfg) - def train_valid_test_data_provider(_train_val_test_num_samples): + def train_valid_test_data_provider(_train_val_test_num_samples): + nonlocal prefetch_loader + loaders = build_train_valid_test_data_loaders(args, topology) + iterators = tuple(iter(loader) if loader is not None else None for loader in loaders) + if not args.mimo_encoder_prefetch or loaders[0] is None: return iterators - train_valid_test_data_provider.is_distributed = True - pretrain( - cfg, - train_valid_test_data_provider, - ModelType.encoder_or_decoder, - mimo_forward_step, - model_provider=None, - skip_model_parallel_init=True, - p2p_communicator=communicator, - pg_collection=topology.schedule_pg_collection, + mimo_model = unwrap_model(captured_model["model"]) + if not mimo_model.role.has_modality_modules: + return iterators + if prefetch_loader is not None: + raise RuntimeError("encoder prefetch loader was already built") + + encoder_module = unwrap_model(mimo_model.modality_submodules[encoder_name]) + prefetch_loader = EncoderPrefetchLoader( + source=iter(loaders[0]), + encoder_name=encoder_name, + feature_producer=partial(prefetch_frozen_features, encoder_module), + depth=args.mimo_encoder_prefetch_depth, + debug=args.mimo_encoder_prefetch_debug, ) - finally: - try: - if topology is not None: - topology.destroy() - finally: - shutdown_distributed() + prefetch_loader.start() + return (prefetch_loader, *iterators[1:]) + + train_valid_test_data_provider.is_distributed = True + pretrain( + cfg, + train_valid_test_data_provider, + ModelType.encoder_or_decoder, + mimo_forward_step, + model_provider=None, + skip_model_parallel_init=True, + p2p_communicator=communicator, + pg_collection=topology.schedule_pg_collection, + ) + + if prefetch_loader is not None: + prefetch_loader.close() + topology.destroy() + shutdown_distributed() if __name__ == "__main__": diff --git a/examples/mimo/scripts/run_hetero_nemotron_20l_mock_train.sh b/examples/mimo/scripts/run_hetero_nemotron_20l_mock_train.sh index a5c759c8e6c..787df50fef8 100755 --- a/examples/mimo/scripts/run_hetero_nemotron_20l_mock_train.sh +++ b/examples/mimo/scripts/run_hetero_nemotron_20l_mock_train.sh @@ -4,8 +4,6 @@ set -euo pipefail -export CUDA_DEVICE_MAX_CONNECTIONS=1 - TRAIN_ITERS=${TRAIN_ITERS:-20} NUM_MICROBATCHES=${NUM_MICROBATCHES:-4} EVAL_INTERVAL=${EVAL_INTERVAL:-1} @@ -96,7 +94,9 @@ uv run --extra ssm python -m torch.distributed.run \ --adam-beta2 0.95 \ --clip-grad 1.0 \ --use-distributed-optimizer \ - --ddp-bucket-size 0 \ + --overlap-grad-reduce \ + --overlap-param-gather \ + --encoder-ddp-overlap \ --train-iters "${TRAIN_ITERS}" \ --eval-interval "${EVAL_INTERVAL}" \ --eval-iters "${EVAL_ITERS}" \ diff --git a/examples/mimo/training/args.py b/examples/mimo/training/args.py index e1d9e2116f7..e62bc31967b 100644 --- a/examples/mimo/training/args.py +++ b/examples/mimo/training/args.py @@ -48,6 +48,14 @@ def add_hetero_grid_args(parser: argparse.ArgumentParser) -> argparse.ArgumentPa "requires --llm-offset 0 so the language grid covers WORLD_SIZE." ), ) + grid.add_argument( + "--encoder-ddp-overlap", + action="store_true", + help=( + "Apply the global grad-reduce and param-gather overlap settings to encoder DDP. " + "Requires every encoder DP rank to execute encoder backward on every microbatch." + ), + ) return parser @@ -56,6 +64,11 @@ def validate_hetero_grid_args(args: argparse.Namespace, world_size: int) -> tupl if args.llm_cp != 1: raise ValueError("hetero MIMO training currently supports CP=1 only") + if getattr(args, "encoder_ddp_overlap", False) and not getattr( + args, "overlap_grad_reduce", False + ): + raise ValueError("--encoder-ddp-overlap requires --overlap-grad-reduce") + # MoE expert count must divide evenly across the language grid's expert parallelism. num_experts = _num_experts(args) if num_experts and num_experts % args.llm_ep != 0: @@ -66,6 +79,8 @@ def validate_hetero_grid_args(args: argparse.Namespace, world_size: int) -> tupl llm_size = args.llm_tp * args.llm_cp * args.llm_pp * args.llm_dp if args.llm_only: + if getattr(args, "encoder_ddp_overlap", False): + raise ValueError("--encoder-ddp-overlap cannot be used with --llm-only") if args.llm_offset != 0: raise ValueError( "--llm-only requires --llm-offset 0 so language ranks cover WORLD_SIZE" diff --git a/examples/mimo/training/batch.py b/examples/mimo/training/batch.py new file mode 100644 index 00000000000..e82af416fae --- /dev/null +++ b/examples/mimo/training/batch.py @@ -0,0 +1,35 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Batch tensor utilities for MIMO training.""" + +from __future__ import annotations + +from collections.abc import Callable +from dataclasses import fields + +import torch + +from megatron.core.packed_seq_params import PackedSeqParams + + +def map_batch_tensors(value, transform: Callable[[torch.Tensor], torch.Tensor]): + """Apply a transform to tensor leaves, including PackedSeqParams fields.""" + if isinstance(value, torch.Tensor): + return transform(value) + if isinstance(value, dict): + return {key: map_batch_tensors(item, transform) for key, item in value.items()} + if isinstance(value, list): + return [map_batch_tensors(item, transform) for item in value] + if isinstance(value, tuple): + return tuple(map_batch_tensors(item, transform) for item in value) + if isinstance(value, PackedSeqParams): + for field in fields(value): + item = getattr(value, field.name) + if isinstance(item, torch.Tensor): + setattr(value, field.name, transform(item)) + return value + + +def move_batch_to_cuda(value): + """Move tensor leaves, including PackedSeqParams tensor fields, to CUDA.""" + return map_batch_tensors(value, lambda tensor: tensor.cuda(non_blocking=True)) diff --git a/examples/mimo/training/data.py b/examples/mimo/training/data.py index 5e1dc187651..9139711662b 100644 --- a/examples/mimo/training/data.py +++ b/examples/mimo/training/data.py @@ -227,7 +227,12 @@ def _build_mock_vlm_dataloader( num_image_tiles=num_image_tiles, ) return DataLoader( - dataset, batch_size=batch_size, shuffle=False, num_workers=0, collate_fn=_collate_mock_batch + dataset, + batch_size=batch_size, + shuffle=False, + num_workers=0, + collate_fn=_collate_mock_batch, + pin_memory=True, ) diff --git a/examples/mimo/training/encoder_prefetch.py b/examples/mimo/training/encoder_prefetch.py new file mode 100644 index 00000000000..c8eb53ddbd1 --- /dev/null +++ b/examples/mimo/training/encoder_prefetch.py @@ -0,0 +1,412 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Bounded read-ahead for a frozen MIMO encoder.""" + +from __future__ import annotations + +import argparse +import logging +import threading +import time +from collections import deque +from collections.abc import Callable + +import torch + +from examples.mimo.training.batch import move_batch_to_cuda + +PREFETCHED_FEATURES_KEY = "_mimo_prefetched_encoder_features" +PROJECTION_TIMER_KEY = "_mimo_encoder_prefetch_projection_timer" + +logger = logging.getLogger(__name__) +_debug_logger = logger.getChild("debug") + + +def add_encoder_prefetch_args(parser: argparse.ArgumentParser) -> argparse.ArgumentParser: + group = parser.add_argument_group("mimo encoder prefetch") + group.add_argument( + "--mimo-encoder-prefetch", + action="store_true", + help="Prefetch completed features from a frozen encoder on encoder ranks.", + ) + group.add_argument( + "--mimo-encoder-prefetch-depth", + type=int, + default=2, + help="Target number of completed encoder-feature batches kept ready.", + ) + group.add_argument( + "--mimo-encoder-prefetch-debug", + action="store_true", + help="Log per-batch encoder-prefetch timing and queue diagnostics.", + ) + return parser + + +def validate_encoder_prefetch_args(args) -> None: + if not args.mimo_encoder_prefetch: + return + if not args.freeze_vit: + raise ValueError("encoder prefetch requires --freeze-vit") + if args.freeze_projection: + raise ValueError("encoder prefetch requires a trainable projection") + for field, label in ( + ("encoder_tp", "TP"), + ("encoder_cp", "CP"), + ("encoder_pp", "PP"), + ("encoder_ep", "EP"), + ): + if getattr(args, field, 1) != 1: + raise ValueError(f"encoder prefetch requires encoder {label}=1") + if args.mimo_encoder_prefetch_depth <= 0: + raise ValueError("encoder prefetch depth must be positive") + if args.rerun_mode != "disabled": + raise ValueError("encoder prefetch does not support rerun modes") + + +def prefetch_frozen_features( + module: torch.nn.Module, encoder_inputs: dict[str, object] +) -> torch.Tensor: + with torch.no_grad(): + return module.combine_embeddings(module.encode(encoder_inputs)) + + +def _record_feature_streams(features: dict[str, torch.Tensor], stream: torch.cuda.Stream) -> None: + for tensor in features.values(): + if tensor.is_cuda: + tensor.record_stream(stream) + + +def _log_producer_debug( + batch_id: int, + data_fetch_ms: float, + encode_start: torch.cuda.Event, + encode_end: torch.cuda.Event, +) -> None: + _debug_logger.info( + "encoder-prefetch-debug producer batch=%d data_fetch_ms=%.3f encode_ms=%.3f", + batch_id, + data_fetch_ms, + encode_start.elapsed_time(encode_end), + ) + + +def _log_consumer_debug( + batch_id: int, ready_at_request: int, depth: int, claimed_pending: bool, wait_start: float +) -> None: + _debug_logger.info( + "encoder-prefetch-debug consumer batch=%d ready_at_request=%d/%d " + "claimed_pending=%d pop_wait_ms=%.3f", + batch_id, + ready_at_request, + depth, + claimed_pending, + (time.perf_counter() - wait_start) * 1000, + ) + + +def _log_encoder_wait_debug( + batch_id: int, start_event: torch.cuda.Event, end_event: torch.cuda.Event +) -> None: + _debug_logger.info( + "encoder-prefetch-debug consumer-wait batch=%d encoder_wait_ms=%.3f", + batch_id, + start_event.elapsed_time(end_event), + ) + + +def _log_projection_debug( + batch_id: int, start_event: torch.cuda.Event, end_event: torch.cuda.Event +) -> None: + _debug_logger.info( + "encoder-prefetch-debug projection batch=%d projection_ms=%.3f", + batch_id, + start_event.elapsed_time(end_event), + ) + + +class _ProjectionTimer: + def __init__(self, loader: EncoderPrefetchLoader, batch_id: int) -> None: + self._loader = loader + self._batch_id = batch_id + self._start_event = torch.cuda.Event(enable_timing=True) + self._end_event = torch.cuda.Event(enable_timing=True) + + def __enter__(self): + self._start_event.record(torch.cuda.current_stream()) + return self + + def __exit__(self, _exc_type, _exc_value, _traceback) -> None: + self._end_event.record(torch.cuda.current_stream()) + self._loader._queue_projection_timing(self._batch_id, self._start_event, self._end_event) + + +class EncoderPrefetchLoader: + """Keep completed encoder features ready without delaying projection.""" + + def __init__( + self, + *, + source, + encoder_name: str, + feature_producer: Callable[[dict[str, object]], torch.Tensor], + depth: int, + stream: torch.cuda.Stream | None = None, + worker_join_timeout_s: float = 30.0, + debug: bool = False, + ) -> None: + """Initialize a bounded encoder-feature prefetch pipeline.""" + if depth <= 0: + raise ValueError("encoder prefetch depth must be positive") + if worker_join_timeout_s <= 0: + raise ValueError("worker_join_timeout_s must be positive") + self._source = iter(source) + self._encoder_name = encoder_name + self._feature_producer = feature_producer + self._depth = depth + self._stream = stream + self._worker_join_timeout_s = worker_join_timeout_s + self._debug = debug + if debug: + _debug_logger.setLevel(logging.INFO) + self._condition = threading.Condition() + self._ready: deque[dict[str, object]] = deque() + self._pending: tuple[dict[str, object], torch.cuda.Event] | None = None + self._encoder_wait_timings: deque[tuple[int, torch.cuda.Event, torch.cuda.Event]] = deque() + self._projection_timings: deque[tuple[int, torch.cuda.Event, torch.cuda.Event]] = deque() + self._produced_batches = 0 + self._consumed_batches = 0 + self._producer_error: BaseException | None = None + self._source_exhausted = False + self._stop = False + self._worker: threading.Thread | None = None + self._device: int | None = None + self._closed = False + + def __iter__(self): + """Return this loader as its own iterator.""" + return self + + def start(self) -> None: + """Initialize CUDA stream state and start the producer thread.""" + with self._condition: + if self._worker is not None: + raise RuntimeError("encoder prefetch loader is already started") + if self._closed: + raise RuntimeError("cannot start a closed encoder prefetch loader") + self._device = torch.cuda.current_device() + if self._stream is None: + self._stream = torch.cuda.Stream() + setup_event = torch.cuda.Event() + setup_event.record(torch.cuda.current_stream()) + self._stream.wait_event(setup_event) + self._worker = threading.Thread( + target=self._producer_main, name=f"mimo-{self._encoder_name}-prefetch", daemon=True + ) + self._worker.start() + + def _producer_main(self) -> None: + """Fetch, encode, and publish batches until exhaustion or shutdown.""" + staged_batch = None + while True: + with self._condition: + # Refill when the completed-feature FIFO has room, or wake for shutdown. + self._condition.wait_for(lambda: self._stop or len(self._ready) < self._depth) + if self._stop: + return + + source_exhausted = False + source_error = None + data_fetch_ms = 0.0 + try: + if staged_batch is None: + batch = next(self._source) + else: + batch = staged_batch + staged_batch = None + item, completion_event, encode_start = self._enqueue_batch(batch) + + with self._condition: + if not self._stop: + self._pending = (item, completion_event) + self._condition.notify_all() + + with self._condition: + should_read_ahead = not self._stop + if should_read_ahead: + # Stage one CPU batch while this GPU encode runs; enqueue it later. + # This advances the source by one batch beyond the trained cursor. + data_fetch_start = time.perf_counter() if self._debug else 0.0 + try: + staged_batch = next(self._source) + except StopIteration: + source_exhausted = True + except BaseException as error: + source_error = error + if self._debug: + data_fetch_ms = (time.perf_counter() - data_fetch_start) * 1000 + completion_event.synchronize() + except StopIteration: + with self._condition: + self._source_exhausted = True + self._condition.notify_all() + return + except BaseException as error: + with self._condition: + if not self._stop: + self._producer_error = error + self._pending = None + self._condition.notify_all() + return + + with self._condition: + if self._stop: + return + if self._pending is not None: + self._ready.append(self._pending[0]) + self._pending = None + batch_id = self._produced_batches + self._produced_batches += 1 + if source_exhausted: + self._source_exhausted = True + if source_error is not None: + self._producer_error = source_error + self._condition.notify_all() + terminate = source_exhausted or source_error is not None + if self._debug: + assert encode_start is not None + _log_producer_debug(batch_id, data_fetch_ms, encode_start, completion_event) + self._drain_encoder_wait_timings() + self._drain_projection_timings() + if terminate: + return + + def _enqueue_batch( + self, batch: dict[str, object] + ) -> tuple[dict[str, object], torch.cuda.Event, torch.cuda.Event | None]: + """Move encoder inputs to CUDA and enqueue their feature computation.""" + if not isinstance(batch, dict): + raise TypeError("encoder prefetch source must return a batch dictionary") + modality_inputs = batch.get("modality_inputs") + if not isinstance(modality_inputs, dict) or self._encoder_name not in modality_inputs: + raise ValueError(f"batch has no inputs for encoder {self._encoder_name!r}") + + with torch.cuda.device(self._device), torch.cuda.stream(self._stream): + # Encoder ranks intentionally retain only fields consumed by their forward step. + output_batch = {"input_ids": batch["input_ids"]} + encoder_inputs = move_batch_to_cuda(modality_inputs[self._encoder_name]) + encode_start = torch.cuda.Event(enable_timing=True) if self._debug else None + if encode_start is not None: + encode_start.record(self._stream) + encoded = self._feature_producer(encoder_inputs) + if not isinstance(encoded, torch.Tensor): + raise TypeError("feature_producer must return one combined tensor") + output_batch[PREFETCHED_FEATURES_KEY] = {self._encoder_name: encoded} + completion_event = torch.cuda.Event(enable_timing=self._debug) + completion_event.record(self._stream) + return output_batch, completion_event, encode_start + + def __next__(self) -> dict[str, object]: + """Return the next feature batch, waiting on a pending encode if needed.""" + if self._worker is None: + raise RuntimeError("encoder prefetch loader must be started before use") + with self._condition: + ready_at_request = len(self._ready) + wait_start = time.perf_counter() if self._debug else 0.0 + self._condition.wait_for( + lambda: self._stop + or self._producer_error is not None + or self._ready + or self._pending is not None + or (self._source_exhausted and not self._ready and self._pending is None) + ) + if self._stop: + raise StopIteration + completion_event = None + if self._ready: + item = self._ready.popleft() + batch_id = self._consumed_batches + self._consumed_batches += 1 + self._condition.notify_all() + elif self._pending is not None: + item, completion_event = self._pending + self._pending = None + batch_id = self._consumed_batches + self._consumed_batches += 1 + self._condition.notify_all() + elif self._producer_error is not None: + raise RuntimeError("encoder prefetch producer failed") from self._producer_error + else: + raise StopIteration + + current_stream = torch.cuda.current_stream() + if completion_event is not None: + wait_start_event = torch.cuda.Event(enable_timing=True) if self._debug else None + wait_end_event = torch.cuda.Event(enable_timing=True) if self._debug else None + if wait_start_event is not None: + wait_start_event.record(current_stream) + current_stream.wait_event(completion_event) + if wait_end_event is not None: + wait_end_event.record(current_stream) + self._queue_encoder_wait_timing(batch_id, wait_start_event, wait_end_event) + _record_feature_streams(item[PREFETCHED_FEATURES_KEY], current_stream) + if self._debug: + item[PROJECTION_TIMER_KEY] = _ProjectionTimer(self, batch_id) + _log_consumer_debug( + batch_id, ready_at_request, self._depth, completion_event is not None, wait_start + ) + return item + + def _queue_encoder_wait_timing( + self, batch_id: int, start_event: torch.cuda.Event, end_event: torch.cuda.Event + ) -> None: + """Queue a consumer encoder-wait timing for asynchronous logging.""" + with self._condition: + self._encoder_wait_timings.append((batch_id, start_event, end_event)) + + def _drain_encoder_wait_timings(self) -> None: + """Log completed consumer encoder-wait timings without synchronizing.""" + ready = [] + with self._condition: + while self._encoder_wait_timings and self._encoder_wait_timings[0][2].query(): + ready.append(self._encoder_wait_timings.popleft()) + for batch_id, start_event, end_event in ready: + _log_encoder_wait_debug(batch_id, start_event, end_event) + + def _queue_projection_timing( + self, batch_id: int, start_event: torch.cuda.Event, end_event: torch.cuda.Event + ) -> None: + """Queue a projection timing for asynchronous logging.""" + with self._condition: + self._projection_timings.append((batch_id, start_event, end_event)) + + def _drain_projection_timings(self) -> None: + """Log completed projection timings without synchronizing.""" + ready = [] + with self._condition: + while self._projection_timings and self._projection_timings[0][2].query(): + ready.append(self._projection_timings.popleft()) + for batch_id, start_event, end_event in ready: + _log_projection_debug(batch_id, start_event, end_event) + + def close(self) -> None: + """Stop the producer and discard buffered batches.""" + with self._condition: + if self._closed: + return + self._closed = True + self._stop = True + self._ready.clear() + self._pending = None + self._condition.notify_all() + worker = self._worker + if worker is not None: + worker.join(timeout=self._worker_join_timeout_s) + if worker.is_alive(): + logger.warning( + "encoder prefetch worker did not stop within %.2f seconds", + self._worker_join_timeout_s, + ) + if self._debug: + self._drain_encoder_wait_timings() + self._drain_projection_timings() diff --git a/examples/mimo/training/grad_sync.py b/examples/mimo/training/grad_sync.py index 9ac6a495aa5..2a06a4b8188 100644 --- a/examples/mimo/training/grad_sync.py +++ b/examples/mimo/training/grad_sync.py @@ -191,3 +191,12 @@ def finalize_grads_func(_model_list, num_tokens, force_all_reduce=False, **_kwar # The schedule always calls grad_scale_func with a Tensor loss; the per-token # mean is applied in finalize_grads_func, so no extra scaling is needed here. mimo_model.config.grad_scale_func = lambda loss: loss + + if getattr(args, "overlap_grad_reduce", False): + assert mimo_model.config.no_sync_func is None, ( + "MIMO overlap owns config.no_sync_func; a second synchronization context " + "cannot be composed safely" + ) + mimo_model.config.no_sync_func = mimo_model.no_sync + if getattr(args, "align_grad_reduce", False): + mimo_model.config.grad_sync_func = mimo_model.start_grad_sync diff --git a/examples/mimo/training/runtime.py b/examples/mimo/training/runtime.py index 4421d12a9e7..6e6427ffd41 100644 --- a/examples/mimo/training/runtime.py +++ b/examples/mimo/training/runtime.py @@ -125,7 +125,9 @@ def wrap_active_modules_with_ddp( [submodule], enc_config, topology.module_pgs[name], - ddp_config=_ddp_config_from_args(args, enable_overlap=False), + ddp_config=_ddp_config_from_args( + args, enable_overlap=getattr(args, "encoder_ddp_overlap", False) + ), data_parallel_random_init=data_parallel_random_init, mixed_precision_wrapper=_EncoderFloat16Module, use_layer_wise_distributed_optimizer=use_layer_wise_distributed_optimizer, diff --git a/examples/mimo/training/step.py b/examples/mimo/training/step.py index ad28ba54189..c9fa387817e 100644 --- a/examples/mimo/training/step.py +++ b/examples/mimo/training/step.py @@ -4,11 +4,13 @@ from __future__ import annotations +from contextlib import nullcontext from functools import partial import torch -from megatron.core.packed_seq_params import PackedSeqParams +from examples.mimo.training.batch import move_batch_to_cuda +from examples.mimo.training.encoder_prefetch import PREFETCHED_FEATURES_KEY, PROJECTION_TIMER_KEY def loss_func(output_tensor: torch.Tensor, *, loss_mask: torch.Tensor): @@ -37,39 +39,28 @@ def loss_func(output_tensor: torch.Tensor, *, loss_mask: torch.Tensor): def mimo_forward_step(data_iterator, model): - """Run a MIMO microbatch for the pipeline schedule. + """Run a raw-input or prefetched-feature MIMO microbatch for the pipeline schedule. On the last pipeline stage, the schedule passes ``output_tensor`` to the returned loss closure. """ batch = next(data_iterator) if data_iterator is not None else {"input_ids": None} - batch = move_batch_to_cuda(batch) - - output_tensor, loss_mask = model(**batch) - return output_tensor, partial(loss_func, loss_mask=loss_mask) - - -def move_batch_to_cuda(value): - """Move tensor leaves, including PackedSeqParams tensor fields, to CUDA.""" - if isinstance(value, torch.Tensor): - return value.cuda(non_blocking=True) - if isinstance(value, dict): - return {key: move_batch_to_cuda(item) for key, item in value.items()} - if isinstance(value, list): - return [move_batch_to_cuda(item) for item in value] - if isinstance(value, tuple): - return tuple(move_batch_to_cuda(item) for item in value) - - if isinstance(value, PackedSeqParams): - for attr in ( - "cu_seqlens_q", - "cu_seqlens_kv", - "cu_seqlens_q_padded", - "cu_seqlens_kv_padded", - "max_seqlen_q", - "max_seqlen_kv", - ): - sub = getattr(value, attr, None) - if isinstance(sub, torch.Tensor) and not sub.is_cuda: - setattr(value, attr, sub.cuda(non_blocking=True)) - return value - return value + prefetched = batch.pop(PREFETCHED_FEATURES_KEY, None) + projection_timer = batch.pop(PROJECTION_TIMER_KEY, None) + + if prefetched is None: + if projection_timer is not None: + raise RuntimeError("encoder prefetch timer has no prefetched features") + batch = move_batch_to_cuda(batch) + output_tensor, loss_mask = model(**batch) + return output_tensor, partial(loss_func, loss_mask=loss_mask) + + if batch.get("modality_inputs"): + raise ValueError("prefetched features cannot be combined with raw modality inputs") + + projection_context = projection_timer if projection_timer is not None else nullcontext() + with projection_context: + output_tensor = model._forward_encoders( + batch.get("input_ids"), modality_inputs=None, input_tensors=prefetched + ) + # Encoder ranks never evaluate the language-model loss closure. + return output_tensor, partial(loss_func, loss_mask=None) diff --git a/examples/post_training/modelopt/quantize.py b/examples/post_training/modelopt/quantize.py index f8d1a3d289c..4e0ef7ed750 100644 --- a/examples/post_training/modelopt/quantize.py +++ b/examples/post_training/modelopt/quantize.py @@ -347,7 +347,7 @@ def get_calib_dataloader( Supports either a local path (.jsonl) or a HuggingFace dataset name. """ - if os.path.isfile(dataset_path_or_name): + if os.path.isfile(dataset_path_or_name) and dataset_path_or_name.endswith(".jsonl"): # Local file print_rank_0(f"Loading calibration dataset from local file: {dataset_path_or_name}") all_texts = [] @@ -539,6 +539,11 @@ def forward_backward_step(model, batch): import_kwargs = {"dtype": import_dtype} if "trust_remote_code" in inspect.signature(import_mcore_gpt_from_hf).parameters: import_kwargs.update({"trust_remote_code": args.trust_remote_code}) + if ( + "moe_router_dtype" in inspect.signature(import_mcore_gpt_from_hf).parameters + and getattr(args, "moe_router_dtype", None) + ): + import_kwargs.update({"moe_router_dtype": args.moe_router_dtype}) import_mcore_gpt_from_hf( unwrapped_model, args.pretrained_model_path, workspace_dir, **import_kwargs ) diff --git a/examples/rl/model_configs/qwen3_30b_a3b_moe.sh b/examples/rl/model_configs/qwen3_30b_a3b_moe.sh index eb55ba35cc6..cb18f8e9885 100644 --- a/examples/rl/model_configs/qwen3_30b_a3b_moe.sh +++ b/examples/rl/model_configs/qwen3_30b_a3b_moe.sh @@ -1,7 +1,8 @@ -#!/bin/bash +#!/bin/bash TP=${TP:-4} PP=${PP:-1} +EP=${EP:-2} NODES_REQUIRED=${NODES_REQUIRED:-1} echo "Using Qwen3-30B-A3B model checkpoint" @@ -33,65 +34,63 @@ ENV_DEPENDENT="\ --grpo-kl-beta $GRPO_KL_BETA \ --langrl-env-config $ENV_CONFIG " - -MODEL_OPTIONS=" ---seq-length $MAX_SEQ_LENGTH \ ---inference-max-seq-length $MAX_SEQ_LENGTH \ ---inference-max-requests $MAX_INFERENCE_BS \ ---pretrained-checkpoint $CHECKPOINT \ ---no-use-tokenizer-model-from-checkpoint-args \ ---seq-length 8192 \ ---inference-max-seq-length 8192 \ ---bf16 \ ---tensor-model-parallel-size $TP \ ---pipeline-model-parallel-size $PP \ ---expert-model-parallel-size $EP \ ---attention-backend flash \ ---transformer-impl transformer_engine \ ---te-rng-tracker \ ---tokenizer-type HuggingFaceTokenizer \ ---tokenizer-model Qwen/Qwen3-30B-A3B \ ---untie-embeddings-and-output-weights \ ---num-layers 48 \ ---hidden-size 2048 \ ---ffn-hidden-size 6144 \ ---num-attention-heads 32 \ ---kv-channels 128 \ ---max-position-embeddings 8192 \ ---group-query-attention \ ---num-query-groups 4 \ ---normalization RMSNorm \ ---norm-epsilon 1e-6 \ ---position-embedding-type rope \ ---rotary-percent 1.0 \ ---rotary-base 1000000 \ ---use-rotary-position-embeddings \ ---swiglu \ ---disable-bias-linear \ ---num-experts 128 \ ---moe-router-topk 8 \ ---moe-ffn-hidden-size 768 \ ---moe-aux-loss-coeff 0.001 \ ---moe-router-load-balancing-type aux_loss \ ---attention-dropout 0.0 \ ---hidden-dropout 0.0 \ ---no-masked-softmax-fusion \ ---attention-softmax-in-fp32 \ ---vocab-size 151936 \ ---make-vocab-size-divisible-by 128 \ ---dist-ckpt-strictness log_unexpected \ ---qk-layernorm \ ---moe-token-dispatcher-type alltoall \ ---moe-layer-freq 1 \ ---optimizer adam \ ---adam-beta1 0.9 \ ---adam-beta2 0.999 \ ---adam-eps 1e-8 \ ---lr 1e-6 \ ---min-lr 1e-7 \ ---lr-warmup-samples 0 \ ---clip-grad 1.0 \ ---weight-decay 0.01 \ ---no-load-optim \ ---ckpt-format torch_dist -" +MODEL_OPTIONS="\ + --seq-length $MAX_SEQ_LENGTH \ + --inference-max-seq-length $MAX_SEQ_LENGTH \ + --inference-max-requests $MAX_INFERENCE_BS \ + --pretrained-checkpoint $CHECKPOINT \ + --no-use-tokenizer-model-from-checkpoint-args \ + --bf16 \ + --tensor-model-parallel-size $TP \ + --pipeline-model-parallel-size $PP \ + --expert-model-parallel-size $EP \ + --attention-backend flash \ + --transformer-impl transformer_engine \ + --te-rng-tracker \ + --tokenizer-type HuggingFaceTokenizer \ + --tokenizer-model Qwen/Qwen3-30B-A3B \ + --tokenizer-hf-include-special-tokens \ + --untie-embeddings-and-output-weights \ + --num-layers 48 \ + --hidden-size 2048 \ + --ffn-hidden-size 6144 \ + --num-attention-heads 32 \ + --kv-channels 128 \ + --max-position-embeddings 8192 \ + --group-query-attention \ + --num-query-groups 4 \ + --normalization RMSNorm \ + --norm-epsilon 1e-6 \ + --position-embedding-type rope \ + --rotary-percent 1.0 \ + --rotary-base 1000000 \ + --use-rotary-position-embeddings \ + --swiglu \ + --disable-bias-linear \ + --num-experts 128 \ + --moe-router-topk 8 \ + --moe-ffn-hidden-size 768 \ + --moe-aux-loss-coeff 0.001 \ + --moe-router-load-balancing-type aux_loss \ + --attention-dropout 0.0 \ + --hidden-dropout 0.0 \ + --no-masked-softmax-fusion \ + --attention-softmax-in-fp32 \ + --vocab-size 151936 \ + --make-vocab-size-divisible-by 128 \ + --dist-ckpt-strictness log_unexpected \ + --qk-layernorm \ + --moe-token-dispatcher-type alltoall \ + --moe-layer-freq 1 \ + --optimizer adam \ + --adam-beta1 0.9 \ + --adam-beta2 0.999 \ + --adam-eps 1e-8 \ + --lr 1e-6 \ + --min-lr 1e-7 \ + --lr-warmup-samples 0 \ + --clip-grad 1.0 \ + --weight-decay 0.01 \ + --no-load-optim \ + --ckpt-format torch_dist \ + " diff --git a/experimental/agent_compose/README.md b/experimental/agent_compose/README.md new file mode 100644 index 00000000000..ca2c8e43e38 --- /dev/null +++ b/experimental/agent_compose/README.md @@ -0,0 +1,45 @@ +# Agent Compose (experimental) + +Agent Compose is an experimental effort to make Megatron-LM development +agentic-native: composing Megatron Core primitives with coding agents, rather +than introducing a new standalone product or training stack. + +This directory is a placeholder that establishes the location and naming for +the upstreamed work. Content will land here incrementally as a series of small, +reviewable PRs. + +## Preview + +The full work-in-progress implementation lives on the `dev` branch under +`experimental/lite/`: + +- https://github.com/NVIDIA/Megatron-LM/tree/dev/experimental/lite + +The preview currently includes: + +- A lightweight runtime API built from small composable primitives. +- Native model implementations with explicit model/runtime protocols. +- Hugging Face safetensors load/export helpers. +- Validation recipes and benchmark examples against Megatron-Core reference + paths (bitwise loss/grad-norm parity on the distributed-optimizer path). +- Skills playbooks that let coding agents extend models and primitives in a + reviewable way. + +## Principles + +- **Compose, don't fork.** Primitives reuse and build from existing Megatron + Core modules wherever appropriate. When a primitive cannot reuse an existing + module and needs a separate implementation, the reason is documented in the + docstring, making gaps explicit and providing input for future Megatron Core + improvements. +- **Reviewable by construction.** Runtime, model, and primitive code are split + into small contracts so agents and humans can make targeted changes without + touching unrelated Megatron subsystems. +- **Core performance.** Changes are validated against Megatron-Core reference + paths for both correctness and speed. + +## Status + +Upstreaming is being scoped: the current work is being evaluated for splitting +into small PRs, after which a timeline will be shared. Until then, please use +the preview branch above. diff --git a/experimental/lite/examples/verl/verl_mlite/compat.py b/experimental/lite/examples/verl/verl_mlite/compat.py index c5786f255b1..a9bbf3ca09f 100644 --- a/experimental/lite/examples/verl/verl_mlite/compat.py +++ b/experimental/lite/examples/verl/verl_mlite/compat.py @@ -18,12 +18,8 @@ from pathlib import Path from typing import Any -_VLLM_ASYNC_SERVER_MODULE = ( - "verl.workers.rollout.vllm_rollout.vllm_async_server" -) -_VLLM_ROLLOUT_CONSUMER_MODULE = ( - "verl.workers.rollout.vllm_rollout.vllm_rollout" -) +_VLLM_ASYNC_SERVER_MODULE = "verl.workers.rollout.vllm_rollout.vllm_async_server" +_VLLM_ROLLOUT_CONSUMER_MODULE = "verl.workers.rollout.vllm_rollout.vllm_rollout" _REGISTERED_HF_CONFIG_TYPES: set[str] = set() _VLLM_IMPORTABLE: bool | None = None @@ -31,6 +27,7 @@ _BUCKETED_SENDER_MODULE = "verl.workers.rollout.vllm_rollout.bucketed_weight_transfer" + def _vllm_importable() -> bool: """Whether ``import vllm`` succeeds in THIS process. @@ -72,11 +69,7 @@ def _register_opaque_hf_config() -> bool: from transformers import AutoConfig, PretrainedConfig - config_cls = type( - "MLiteOpaqueConfig", - (PretrainedConfig,), - {"model_type": model_type}, - ) + config_cls = type("MLiteOpaqueConfig", (PretrainedConfig,), {"model_type": model_type}) try: AutoConfig.register(model_type, config_cls) except ValueError: @@ -121,10 +114,7 @@ def _install_vllm_thin_finder() -> bool: # own (container/SM90) vllm. if _vllm_site_ld_library_path(site): return False - if any( - getattr(finder, "_verl_mlite_vllm_thin_finder", False) - for finder in sys.meta_path - ): + if any(getattr(finder, "_verl_mlite_vllm_thin_finder", False) for finder in sys.meta_path): return False sys.meta_path.insert(0, _VllmThinFinder(site)) return True @@ -156,9 +146,7 @@ def _patch_transformers_vision2seq_alias() -> bool: except Exception: # pragma: no cover - transformers internals moved lazy_cls = None - if lazy_cls is not None and not getattr( - lazy_cls, "_mlite_vision2seq_patched", False - ): + if lazy_cls is not None and not getattr(lazy_cls, "_mlite_vision2seq_patched", False): _orig_getattr = lazy_cls.__getattr__ def _getattr_with_vision2seq_alias(self, name): @@ -241,10 +229,7 @@ def _vllm_server_profile_env() -> dict[str, str]: pythonpath_entries = _vllm_site_pythonpath_prefixes(site) + [site] if pythonpath: pythonpath_entries.append(pythonpath) - result = { - "PYTHONPATH": os.pathsep.join(pythonpath_entries), - "PYTHONNOUSERSITE": "1", - } + result = {"PYTHONPATH": os.pathsep.join(pythonpath_entries), "PYTHONNOUSERSITE": "1"} ld_library_entries = _vllm_site_ld_library_path(site) if ld_library_entries: existing_ld = os.environ.get("LD_LIBRARY_PATH", "").strip() @@ -288,11 +273,7 @@ def _patch_verl_vllm_headless_api_server_count() -> bool: server_module = importlib.import_module(_VLLM_ASYNC_SERVER_MODULE) original_run_headless = server_module.run_headless - if getattr( - original_run_headless, - "_verl_mlite_api_server_count_patch", - False, - ): + if getattr(original_run_headless, "_verl_mlite_api_server_count_patch", False): return False @wraps(original_run_headless) @@ -359,11 +340,10 @@ def _patch_verl_vllm_device_uuid() -> bool: if getattr(original_get_device_uuid, "_verl_mlite_visible_device_patch", False): patched_get_device_uuid = original_get_device_uuid else: + @wraps(original_get_device_uuid) def patched_get_device_uuid(device_id: int) -> str: - return original_get_device_uuid( - _normalize_vllm_visible_device_id(device_id) - ) + return original_get_device_uuid(_normalize_vllm_visible_device_id(device_id)) patched_get_device_uuid._verl_mlite_visible_device_patch = True utils.get_device_uuid = patched_get_device_uuid @@ -485,14 +465,10 @@ def raw_binding(name: str) -> tuple[Any, str]: return missing, "absent" alias, alias_source = raw_binding("AutoModelForVision2Seq") - replacement, replacement_source = raw_binding( - "AutoModelForImageTextToText" - ) + replacement, replacement_source = raw_binding("AutoModelForImageTextToText") payload = { "alias_is_replacement": ( - alias is not missing - and replacement is not missing - and alias is replacement + alias is not missing and replacement is not missing and alias is replacement ), "alias_source": alias_source, "changed": result, @@ -504,10 +480,7 @@ def raw_binding(name: str) -> tuple[Any, str]: "transformers_id": id(transformers) if transformers is not None else None, "transformers_loaded": transformers is not None, } - sys.stderr.write( - "VERL_MLITE_RUNTIME_PATCH_TRACE " - f"{json.dumps(payload, sort_keys=True)}\n" - ) + sys.stderr.write("VERL_MLITE_RUNTIME_PATCH_TRACE " f"{json.dumps(payload, sort_keys=True)}\n") sys.stderr.flush() @@ -628,6 +601,7 @@ def _recreate_dense_fp8_linear_params(model) -> int: the single post-load process_weights_after_loading matches cold load bit-for-bit (see module block above). Returns the count recreated.""" import torch + try: from vllm.model_executor.layers.linear import LinearBase from vllm.model_executor.layers.quantization.fp8 import Fp8LinearMethod @@ -635,8 +609,10 @@ def _recreate_dense_fp8_linear_params(model) -> int: return 0 recreated = 0 for _name, layer in model.named_modules(): - if not (isinstance(layer, LinearBase) - and isinstance(getattr(layer, "quant_method", None), Fp8LinearMethod)): + if not ( + isinstance(layer, LinearBase) + and isinstance(getattr(layer, "quant_method", None), Fp8LinearMethod) + ): continue qm = layer.quant_method if not getattr(qm, "block_quant", False): @@ -659,9 +635,7 @@ def _recreate_dense_fp8_linear_params(model) -> int: except Exception as _re: sys.stderr.write(f"VERL_MLITE_DENSE_RECREATE_SKIP {_name}: {_re!r}\n") sys.stderr.flush() - raise RuntimeError( - f"failed to recreate dense FP8 parameters for {_name}" - ) from _re + raise RuntimeError(f"failed to recreate dense FP8 parameters for {_name}") from _re return recreated @@ -696,7 +670,8 @@ def prepare_quanted_weights_for_loading(model_runner, *args, **kwargs): if n: sys.stderr.write( f"VERL_MLITE_DENSE_RECREATE recreated {n} dense FP8 linear " - "param set(s) to checkpoint layout before resync load\n") + "param set(s) to checkpoint layout before resync load\n" + ) sys.stderr.flush() except Exception as exc: sys.stderr.write(f"VERL_MLITE_DENSE_RECREATE error: {exc!r}\n") @@ -782,9 +757,7 @@ def _patch_verl_dsv4_native_layerwise_reload() -> bool: try: fp8_utils = importlib.import_module("verl.utils.vllm.vllm_fp8_utils") dsv4_utils = importlib.import_module("verl.utils.vllm.vllm_dsv4_fp8_utils") - rollout_utils = importlib.import_module( - "verl.workers.rollout.vllm_rollout.utils" - ) + rollout_utils = importlib.import_module("verl.workers.rollout.vllm_rollout.utils") except Exception: return False original_prepare = getattr(fp8_utils, "prepare_quanted_weights_for_loading", None) @@ -808,15 +781,11 @@ def prepare_quanted_weights_for_loading(model_runner, *args, **kwargs): # updated directly by VERL's buffer path (or restored below). Keeping # them out of the meta restore prevents ``copy_`` into a meta buffer # from silently discarding router state between IPC buckets. - SKIP_TENSORS.update( - {"tid2eid", "expert_bias", "e_score_correction_bias", "attn_sink"} - ) + SKIP_TENSORS.update({"tid2eid", "expert_bias", "e_score_correction_bias", "attn_sink"}) with set_current_vllm_config(model_runner.vllm_config): initialize_layerwise_reload(model) model._verl_mlite_ds4_layerwise_reload_active = True - sys.stderr.write( - "VERL_MLITE_DSV4_LAYERWISE_RELOAD initialized native vLLM reload\n" - ) + sys.stderr.write("VERL_MLITE_DSV4_LAYERWISE_RELOAD initialized native vLLM reload\n") sys.stderr.flush() return _DSV4_LAYERWISE_RELOAD_STATE @@ -830,8 +799,7 @@ def process_quanted_weights_after_loading(model_runner, reload_state): try: with set_current_vllm_config(model_runner.vllm_config): finalize_layerwise_processing( - model_runner.model, - model_runner.vllm_config.model_config, + model_runner.model, model_runner.vllm_config.model_config ) finally: model_runner.model._verl_mlite_ds4_layerwise_reload_active = False @@ -878,9 +846,7 @@ def load_quanted_weights(weights, model_runner, *args, **kwargs): if records is None: records = [] model._verl_mlite_weight_fingerprint = records - records.extend( - tensor_fingerprint_record(name, tensor) for name, tensor in weights - ) + records.extend(tensor_fingerprint_record(name, tensor) for name, tensor in weights) # Layerwise reload may retain loader arguments across multiple IPC # callbacks. The receiver reuses its communication buffer after # each callback, so persist every tensor until its logical layer is @@ -912,9 +878,7 @@ def _patch_verl_dsv4_fp8_process_weights() -> bool: if not _vllm_importable(): return False try: - utils = importlib.import_module( - "verl.workers.rollout.vllm_rollout.utils" - ) + utils = importlib.import_module("verl.workers.rollout.vllm_rollout.utils") except Exception: return False ext_cls = getattr(utils, "vLLMColocateWorkerExtension", None) @@ -998,11 +962,7 @@ def next_bucket(self, staging): break layer_key = self._layer_cluster_key(name) - if ( - bucket_meta - and bucket_layer_key is not None - and layer_key != bucket_layer_key - ): + if bucket_meta and bucket_layer_key is not None and layer_key != bucket_layer_key: self._pending = (name, weight) break @@ -1120,7 +1080,9 @@ def produce(): except queue.Empty: if worker_future.done(): worker_future.result() - raise RuntimeError("MLite weight prefetch stopped without a terminal result") + raise RuntimeError( + "MLite weight prefetch stopped without a terminal result" + ) continue kind, metadata_or_name, direct_weight, used_bytes, ready, is_last, held_slot = ( @@ -1149,9 +1111,7 @@ def produce(): free_slots.put_nowait(held_slot) held_slot = None - self.socket.send_pyobj( - {"bucket_meta": metadata_or_name, "is_last": is_last} - ) + self.socket.send_pyobj({"bucket_meta": metadata_or_name, "is_last": is_last}) self.socket.recv() if is_last: break @@ -1172,12 +1132,7 @@ def produce(): def _weight_sync_probe_enabled() -> bool: - return os.getenv("MLITE_WEIGHT_SYNC_PROBE", "").strip().lower() in { - "1", - "true", - "yes", - "on", - } + return os.getenv("MLITE_WEIGHT_SYNC_PROBE", "").strip().lower() in {"1", "true", "yes", "on"} def _weight_sync_fingerprint_enabled() -> bool: diff --git a/experimental/lite/examples/verl/verl_mlite/engine/config.py b/experimental/lite/examples/verl/verl_mlite/engine/config.py index d5f8b73c284..79bac426e46 100644 --- a/experimental/lite/examples/verl/verl_mlite/engine/config.py +++ b/experimental/lite/examples/verl/verl_mlite/engine/config.py @@ -48,11 +48,7 @@ def __post_init__(self) -> None: if self.resync_format is not None: from megatron.lite.runtime.contracts.weights import ResyncFormat - object.__setattr__( - self, - "resync_format", - ResyncFormat.parse(self.resync_format).value, - ) + object.__setattr__(self, "resync_format", ResyncFormat.parse(self.resync_format).value) if not isinstance(self.resync_config, Mapping): raise TypeError("resync_config must be a mapping") object.__setattr__(self, "resync_config", dict(self.resync_config)) diff --git a/experimental/lite/megatron/lite/model/qwen3_moe/lite/protocol.py b/experimental/lite/megatron/lite/model/qwen3_moe/lite/protocol.py index 0e3220639bc..79d52f286dc 100644 --- a/experimental/lite/megatron/lite/model/qwen3_moe/lite/protocol.py +++ b/experimental/lite/megatron/lite/model/qwen3_moe/lite/protocol.py @@ -26,6 +26,7 @@ import torch import torch.nn as nn + from megatron.lite.model.protocol_utils import ( add_cross_entropy_fusion, add_loss_context_kwargs, @@ -324,10 +325,7 @@ def export_hf_weights( def save_hf_weights( - chunks: list[nn.Module], - path: str, - model_cfg: Qwen3MoEConfig, - ps: ParallelState, + chunks: list[nn.Module], path: str, model_cfg: Qwen3MoEConfig, ps: ParallelState ) -> None: from megatron.lite.model.qwen3_moe.lite.checkpoint import save_hf_weights as _save diff --git a/experimental/lite/megatron/lite/primitive/parallel/pipeline.py b/experimental/lite/megatron/lite/primitive/parallel/pipeline.py index 8cbc775867f..f665d2091cc 100644 --- a/experimental/lite/megatron/lite/primitive/parallel/pipeline.py +++ b/experimental/lite/megatron/lite/primitive/parallel/pipeline.py @@ -290,13 +290,7 @@ def _p2p(send_fwd=None, send_bwd=None, recv_fwd=False, recv_bwd=False): # Megatron dynamic shape exchange: the recv buffer is sized from the shape # the sender transmits, so no per-mb shape has to be tracked here. return _send_recv_pipeline( - send_fwd, - send_bwd, - recv_fwd, - recv_bwd, - ps, - tensor_shape, - dynamic_shape=True, + send_fwd, send_bwd, recv_fwd, recv_bwd, ps, tensor_shape, dynamic_shape=True ) # ── Warmup: pure forward passes ── @@ -574,12 +568,16 @@ def _send_recv_pipeline( ops.append(dist.P2POp(dist.isend, t, ps.pp_next_rank, p2p_group)) if recv_fwd: if dynamic_shape: - fwd_buf = torch.empty(recv_fwd_shape, dtype=_PIPELINE_TENSOR_DTYPE, device=_pipeline_device()) + fwd_buf = torch.empty( + recv_fwd_shape, dtype=_PIPELINE_TENSOR_DTYPE, device=_pipeline_device() + ) else: fwd_buf = ( fwd_recv_buf if fwd_recv_buf is not None - else torch.empty(tensor_shape, dtype=_PIPELINE_TENSOR_DTYPE, device=_pipeline_device()) + else torch.empty( + tensor_shape, dtype=_PIPELINE_TENSOR_DTYPE, device=_pipeline_device() + ) ) ops.append(dist.P2POp(dist.irecv, fwd_buf, ps.pp_prev_rank, p2p_group)) if send_bwd is not None: @@ -587,12 +585,16 @@ def _send_recv_pipeline( ops.append(dist.P2POp(dist.isend, t, ps.pp_prev_rank, p2p_group)) if recv_bwd: if dynamic_shape: - bwd_buf = torch.empty(recv_bwd_shape, dtype=_PIPELINE_TENSOR_DTYPE, device=_pipeline_device()) + bwd_buf = torch.empty( + recv_bwd_shape, dtype=_PIPELINE_TENSOR_DTYPE, device=_pipeline_device() + ) else: bwd_buf = ( bwd_recv_buf if bwd_recv_buf is not None - else torch.empty(tensor_shape, dtype=_PIPELINE_TENSOR_DTYPE, device=_pipeline_device()) + else torch.empty( + tensor_shape, dtype=_PIPELINE_TENSOR_DTYPE, device=_pipeline_device() + ) ) ops.append(dist.P2POp(dist.irecv, bwd_buf, ps.pp_next_rank, p2p_group)) diff --git a/experimental/lite/tests/unit/model/test_qwen35_export.py b/experimental/lite/tests/unit/model/test_qwen35_export.py index 1e602f48c7f..84683b33dfd 100644 --- a/experimental/lite/tests/unit/model/test_qwen35_export.py +++ b/experimental/lite/tests/unit/model/test_qwen35_export.py @@ -162,10 +162,7 @@ def __init__(self) -> None: exported = dict( export_hf_weights( - TinyQwen35Module(), - _tiny_config(), - _single_rank_parallel_state(), - cpu=False, + TinyQwen35Module(), _tiny_config(), _single_rank_parallel_state(), cpu=False ) ) @@ -184,9 +181,9 @@ def __init__(self, config: Qwen35Config) -> None: layer.moe.experts = nn.Module() layer.moe.experts.fc1 = nn.Module() for local_idx in range(config.num_experts // 2): - tensor = torch.arange( - rows * config.hidden_size, dtype=torch.bfloat16 - ).reshape(rows, config.hidden_size) + tensor = torch.arange(rows * config.hidden_size, dtype=torch.bfloat16).reshape( + rows, config.hidden_size + ) tensor = tensor + layer_idx * 10000 + local_idx * 1000 layer.moe.experts.fc1.register_parameter( f"weight{local_idx}", nn.Parameter(tensor) @@ -215,8 +212,7 @@ def fake_all_gather(output, tensor, group=None): output[tensor.numel() :].copy_(tensor + 2000) monkeypatch.setattr( - "megatron.lite.primitive.ckpt.hf_weights.dist.all_gather_into_tensor", - fake_all_gather, + "megatron.lite.primitive.ckpt.hf_weights.dist.all_gather_into_tensor", fake_all_gather ) exported = dict(export_hf_weights(model, cfg, ps)) @@ -233,12 +229,9 @@ def fake_all_gather(output, tensor, group=None): getattr(model.layers[1].moe.experts.fc1, f"weight{i}").detach() for i in range(cfg.num_experts // ps.ep_size) ] - layer1_expected = torch.stack( - layer1_tensors + [tensor + 2000 for tensor in layer1_tensors] - ) + layer1_expected = torch.stack(layer1_tensors + [tensor + 2000 for tensor in layer1_tensors]) assert torch.equal( - exported["model.language_model.layers.1.mlp.experts.gate_up_proj"], - layer1_expected, + exported["model.language_model.layers.1.mlp.experts.gate_up_proj"], layer1_expected ) @@ -314,8 +307,7 @@ def fake_all_gather(output, tensor, group=None): monkeypatch.setattr("megatron.lite.primitive.ckpt.hf_weights.dist.is_initialized", lambda: True) monkeypatch.setattr("megatron.lite.primitive.ckpt.hf_weights.dist.get_rank", lambda: 1) monkeypatch.setattr( - "megatron.lite.primitive.ckpt.hf_weights.dist.all_gather_into_tensor", - fake_all_gather, + "megatron.lite.primitive.ckpt.hf_weights.dist.all_gather_into_tensor", fake_all_gather ) exported = list(export_hf_weights(TinyQwen35Module(cfg), cfg, ps, rank0_only=True)) @@ -496,8 +488,7 @@ def fake_all_gather(output, tensor, group=None): output[tensor.numel() :].copy_(shards[1].view(-1)) monkeypatch.setattr( - "megatron.lite.primitive.ckpt.hf_weights.dist.all_gather_into_tensor", - fake_all_gather, + "megatron.lite.primitive.ckpt.hf_weights.dist.all_gather_into_tensor", fake_all_gather ) exported = dict(export_hf_weights(TinyQwen35Module(shards[0]), cfg, ps)) diff --git a/experimental/lite/tests/unit/primitive/test_fsdp2_offload_gpu.py b/experimental/lite/tests/unit/primitive/test_fsdp2_offload_gpu.py index e1c21179807..7a2b21ec5b7 100644 --- a/experimental/lite/tests/unit/primitive/test_fsdp2_offload_gpu.py +++ b/experimental/lite/tests/unit/primitive/test_fsdp2_offload_gpu.py @@ -79,9 +79,7 @@ def forward(self, x): class MemoryStressExperts(nn.Module): def __init__(self, numel: int): super().__init__() - self.weight = nn.Parameter( - torch.empty(numel, device="cuda", dtype=torch.bfloat16) - ) + self.weight = nn.Parameter(torch.empty(numel, device="cuda", dtype=torch.bfloat16)) def forward(self, x): return x + self.weight[0] * 0 @@ -90,9 +88,7 @@ def forward(self, x): class MemoryStressMoEUnit(nn.Module): def __init__(self, expert_numel: int): super().__init__() - self.dense_scale = nn.Parameter( - torch.ones(1, device="cuda", dtype=torch.bfloat16) - ) + self.dense_scale = nn.Parameter(torch.ones(1, device="cuda", dtype=torch.bfloat16)) self.experts = MemoryStressExperts(expert_numel) def forward(self, x): @@ -102,9 +98,7 @@ def forward(self, x): class MemoryStressMoEModel(nn.Module): def __init__(self, expert_numel: int, num_units: int): super().__init__() - self.units = nn.ModuleList( - MemoryStressMoEUnit(expert_numel) for _ in range(num_units) - ) + self.units = nn.ModuleList(MemoryStressMoEUnit(expert_numel) for _ in range(num_units)) def forward(self, x): for unit in self.units: @@ -243,9 +237,7 @@ def test_fsdp2_pp_edp_reshard_and_offload_roundtrip_eight_gpus(): torch.manual_seed(1234) model = TinyPipelineMoEModel().cuda().to(dtype=torch.bfloat16) expert_global_numels = { - name: param.numel() - for name, param in model.named_parameters() - if ".experts." in name + name: param.numel() for name, param in model.named_parameters() if ".experts." in name } optimizer = build_fsdp2_training_optimizer( [model], @@ -280,16 +272,13 @@ def assert_experts_are_local_shards(device: str) -> None: def assert_grad_devices(device: str) -> None: grads = [param.grad for param in model.parameters()] assert all(grad is not None for grad in grads) - assert {to_local_tensor(grad).device.type for grad in grads if grad is not None} == { - device - } + assert {to_local_tensor(grad).device.type for grad in grads if grad is not None} == {device} def check_expert_shard_after_forward(_module, _inputs, _output) -> None: assert_experts_are_local_shards("cuda") hooks = [ - unit.experts.register_forward_hook(check_expert_shard_after_forward) - for unit in model.units + unit.experts.register_forward_hook(check_expert_shard_after_forward) for unit in model.units ] def train_backward(seed: int) -> None: @@ -362,8 +351,7 @@ def check_stress_expert_shard(_module, _inputs, _output) -> None: stress_shard_checks.append(True) stress_hooks = [ - unit.experts.register_forward_hook(check_stress_expert_shard) - for unit in stress_model.units + unit.experts.register_forward_hook(check_stress_expert_shard) for unit in stress_model.units ] device_bytes = torch.cuda.get_device_properties(torch.cuda.current_device()).total_memory memory_limit_bytes = resident_bytes + 3 * expert_bytes @@ -371,9 +359,7 @@ def check_stress_expert_shard(_module, _inputs, _output) -> None: torch.cuda.set_per_process_memory_fraction(memory_fraction) try: with torch.no_grad(): - stress_output = stress_model( - torch.ones(1, device="cuda", dtype=torch.bfloat16) - ) + stress_output = stress_model(torch.ones(1, device="cuda", dtype=torch.bfloat16)) torch.cuda.synchronize() assert torch.cuda.memory_allocated() <= resident_bytes + expert_bytes assert len(stress_shard_checks) == len(stress_model.units) @@ -397,8 +383,8 @@ def check_stress_expert_shard(_module, _inputs, _output) -> None: # barrier: parameters are intentionally left materialized after forward, # then runtime.to(..., "cpu") must reshard them before Module.to() sees # their DTensor state. - materialized_model = TinyPipelineMoEModel(hidden_size=512, num_units=1).cuda().to( - dtype=torch.bfloat16 + materialized_model = ( + TinyPipelineMoEModel(hidden_size=512, num_units=1).cuda().to(dtype=torch.bfloat16) ) materialized_expert_numels = { name: param.numel() diff --git a/experimental/lite/tests/unit/runtime/test_runtime_backend_unit.py b/experimental/lite/tests/unit/runtime/test_runtime_backend_unit.py index fc41b6218aa..9571838b88d 100644 --- a/experimental/lite/tests/unit/runtime/test_runtime_backend_unit.py +++ b/experimental/lite/tests/unit/runtime/test_runtime_backend_unit.py @@ -242,28 +242,19 @@ def load_state_to_device(self): import megatron.lite.runtime.megatron_utils as megatron_utils monkeypatch.setattr( - megatron_utils, - "offload_model_to_cpu", - lambda chunks: events.append("offload-model"), + megatron_utils, "offload_model_to_cpu", lambda chunks: events.append("offload-model") ) monkeypatch.setattr( - megatron_utils, - "load_model_to_gpu", - lambda chunks, load_grad: events.append("load-model"), + megatron_utils, "load_model_to_gpu", lambda chunks, load_grad: events.append("load-model") ) monkeypatch.setattr(torch.cuda, "is_available", lambda: True) monkeypatch.setattr(torch.cuda, "synchronize", lambda: events.append("synchronize")) monkeypatch.setattr( - "megatron.lite.runtime.backends.mlite.runtime.gc.collect", - lambda: events.append("collect"), + "megatron.lite.runtime.backends.mlite.runtime.gc.collect", lambda: events.append("collect") ) monkeypatch.setattr(torch.cuda, "empty_cache", lambda: events.append("empty-cache")) chunk = Chunk() - handle = ModelHandle( - model=chunk, - optimizer=Optimizer(), - _extras={"model_chunks": [chunk]}, - ) + handle = ModelHandle(model=chunk, optimizer=Optimizer(), _extras={"model_chunks": [chunk]}) runtime = MegatronLiteRuntime.__new__(MegatronLiteRuntime) runtime.to(handle, "cpu", model=True, optimizer=False, grad=True) @@ -291,15 +282,11 @@ def release_export_scratch(self): import megatron.lite.runtime.megatron_utils as megatron_utils monkeypatch.setattr( - megatron_utils, - "offload_model_to_cpu", - lambda chunks: events.append("offload-model"), + megatron_utils, "offload_model_to_cpu", lambda chunks: events.append("offload-model") ) chunk = Chunk() handle = ModelHandle( - model=chunk, - optimizer=HookedOptimizer(), - _extras={"model_chunks": [chunk]}, + model=chunk, optimizer=HookedOptimizer(), _extras={"model_chunks": [chunk]} ) MegatronLiteRuntime.__new__(MegatronLiteRuntime).to( diff --git a/experimental/lite/tests/unit/verl/test_mlite_engine_checkpoint.py b/experimental/lite/tests/unit/verl/test_mlite_engine_checkpoint.py index f2104d0fc71..eab479b768b 100644 --- a/experimental/lite/tests/unit/verl/test_mlite_engine_checkpoint.py +++ b/experimental/lite/tests/unit/verl/test_mlite_engine_checkpoint.py @@ -210,9 +210,7 @@ def test_hf_model_save_fails_loudly_when_model_config_is_missing(tmp_path): def test_hf_model_only_save_uses_protocol_and_writes_hf_metadata(tmp_path, monkeypatch): - engine, module, *_ = _initialized_engine( - checkpoint_config={"save_contents": ["hf_model"]} - ) + engine, module, *_ = _initialized_engine(checkpoint_config={"save_contents": ["hf_model"]}) model_cfg = object() chunks = [module, object()] export_calls = [] @@ -235,25 +233,17 @@ def save_pretrained(self, path): { "model_cfg": model_cfg, "model_chunks": chunks, - "protocol": SimpleNamespace( - save_hf_weights=lambda *args: export_calls.append(args) - ), + "protocol": SimpleNamespace(save_hf_weights=lambda *args: export_calls.append(args)), } ) monkeypatch.setattr( "verl_mlite.engine.mlite_engine.save_training_checkpoint", - lambda *args, **kwargs: pytest.fail( - "hf_model-only save wrote a native checkpoint" - ), + lambda *args, **kwargs: pytest.fail("hf_model-only save wrote a native checkpoint"), ) engine.save_checkpoint(str(tmp_path), global_step=1) hf_path = str(tmp_path / "huggingface") assert export_calls == [(chunks, hf_path, model_cfg, engine.handle._parallel_state)] - assert metadata_calls == [ - ("config", hf_path), - ("tokenizer", hf_path), - ("processor", hf_path), - ] + assert metadata_calls == [("config", hf_path), ("tokenizer", hf_path), ("processor", hf_path)] assert hf_config.auto_map == {"AutoModel": "modeling.Model"} diff --git a/megatron/core/datasets/indexed_dataset.py b/megatron/core/datasets/indexed_dataset.py index 76de4cca8d2..ce216c8b855 100644 --- a/megatron/core/datasets/indexed_dataset.py +++ b/megatron/core/datasets/indexed_dataset.py @@ -39,7 +39,7 @@ is_object_storage_path, parse_s3_path, ) -from megatron.core.msc_utils import MultiStorageClientFeature +from megatron.core.msc_utils import MultiStorageClientFeature, maybe_msc from megatron.core.utils import log_single_rank logger = logging.getLogger(__name__) @@ -138,11 +138,7 @@ def __enter__(self) -> "_IndexWriter": Returns: _IndexWriter: The instance """ - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - self.idx_writer = msc.open(self.idx_path, "wb") - else: - self.idx_writer = open(self.idx_path, "wb") + self.idx_writer = maybe_msc.open(self.idx_path, "wb") # fixed, vestigial practice self.idx_writer.write(_INDEX_HEADER) # fixed, vestigial practice @@ -394,11 +390,7 @@ class _MMapBinReader(_BinReader): """ def __init__(self, bin_path: str) -> None: - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - self._bin_file_reader = msc.open(bin_path, mode="rb") - else: - self._bin_file_reader = open(bin_path, mode="rb") + self._bin_file_reader = maybe_msc.open(bin_path, mode="rb") self._bin_buffer_mmap = numpy.memmap(self._bin_file_reader, mode="r", order="C") self._bin_buffer = memoryview(self._bin_buffer_mmap.data) @@ -462,15 +454,9 @@ def read(self, dtype: Type[numpy.number], count: int, offset: int) -> numpy.ndar def _read(): """Helper method to read `count` bytes from self._bin_path at provided offset.""" sequence = numpy.empty(count, dtype=dtype) - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - with msc.open(self._bin_path, mode="rb", buffering=0) as bin_buffer_file: - bin_buffer_file.seek(offset) - bin_buffer_file.readinto(sequence) - else: - with open(self._bin_path, mode="rb", buffering=0) as bin_buffer_file: - bin_buffer_file.seek(offset) - bin_buffer_file.readinto(sequence) + with maybe_msc.open(self._bin_path, mode="rb", buffering=0) as bin_buffer_file: + bin_buffer_file.seek(offset) + bin_buffer_file.readinto(sequence) return sequence sleep_duration = self.sleep_duration_start @@ -948,13 +934,7 @@ class IndexedDatasetBuilder(object): def __init__( self, bin_path: str, dtype: Type[numpy.number] = numpy.int32, multimodal: bool = False ) -> None: - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - self._open = msc.open - else: - self._open = open - - self.data_file = self._open(bin_path, "wb") + self.data_file = maybe_msc.open(bin_path, "wb") self.dtype = dtype self.multimodal = multimodal @@ -1023,7 +1003,7 @@ def add_index(self, path_prefix: str) -> None: gc.collect() # Concatenate data - with self._open(get_bin_path(path_prefix), "rb") as f: + with maybe_msc.open(get_bin_path(path_prefix), "rb") as f: shutil.copyfileobj(f, self.data_file) def finalize(self, idx_path: str) -> None: diff --git a/megatron/core/dist_checkpointing/core.py b/megatron/core/dist_checkpointing/core.py index c601d0f5ce9..cdb244dbb8d 100644 --- a/megatron/core/dist_checkpointing/core.py +++ b/megatron/core/dist_checkpointing/core.py @@ -8,7 +8,7 @@ from dataclasses import asdict, dataclass from typing import Optional -from megatron.core.msc_utils import MultiStorageClientFeature +from megatron.core.msc_utils import maybe_msc CONFIG_FNAME = 'metadata.json' @@ -57,17 +57,10 @@ def maybe_load_config(checkpoint_dir: str) -> Optional[CheckpointingConfig]: """ config_path = os.path.join(checkpoint_dir, CONFIG_FNAME) if checkpoint_dir: - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - if not msc.os.path.exists(config_path): - return None - with msc.open(config_path) as f: - config_dict = json.load(f) - else: - if not os.path.exists(config_path): - return None - with open(config_path) as f: - config_dict = json.load(f) + if not maybe_msc.os.path.exists(config_path): + return None + with maybe_msc.open(config_path) as f: + config_dict = json.load(f) known_fields = {f.name for f in dataclasses.fields(CheckpointingConfig)} return CheckpointingConfig(**{k: v for k, v in config_dict.items() if k in known_fields}) return None @@ -84,10 +77,5 @@ def save_config(config: CheckpointingConfig, checkpoint_dir: str): None """ config_path = os.path.join(checkpoint_dir, CONFIG_FNAME) - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - with msc.open(config_path, 'w') as f: - json.dump(asdict(config), f) - else: - with open(config_path, 'w') as f: - json.dump(asdict(config), f) + with maybe_msc.open(config_path, 'w') as f: + json.dump(asdict(config), f) diff --git a/megatron/core/dist_checkpointing/gpt_checkpoint_interop.py b/megatron/core/dist_checkpointing/gpt_checkpoint_interop.py new file mode 100644 index 00000000000..47baffda9df --- /dev/null +++ b/megatron/core/dist_checkpointing/gpt_checkpoint_interop.py @@ -0,0 +1,362 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Load GPT (pure transformer) distributed checkpoints into HybridModel runs. + +A GPTModel decoder layer packs self-attention and an MLP into a single +``TransformerLayer``, so a GPT checkpoint with ``L`` layers stores both +sub-modules under ``decoder.layers..``. HybridModel gives every sub-block +its own layer: in a pattern such as ``M*-M*-`` each GPT layer corresponds to +one attention ('*') position and one MLP ('-' dense or 'E' MoE) position, +while SSM ('M') positions have no GPT counterpart. + +Rather than rewriting the checkpoint on disk, the hybrid run's own sharded +state dict is retargeted at load time: + +* attention and MLP entries are rewritten to the GPT checkpoint's canonical + homogeneous-layer format: the layer index is dropped from the storage + ``key`` (``decoder.layers..mlp...`` -> ``decoder.layers.mlp...``) and + the matching GPT layer index becomes a prepended sharding axis, exactly + mirroring ``TransformerBlock.sharded_state_dict`` with + ``non_homogeneous_layers=False`` (the format GPTModel training saves); +* ``decoder.final_norm`` is pointed at GPT's ``decoder.final_layernorm``; +* HybridModel's empty ``output_layer._extra_state`` entry stays local because + GPT checkpoints intentionally omit that backward-compatibility key; +* entries of layers without a GPT counterpart are wrapped in + ``LocalNonpersistentObject`` so no storage read is attempted and the + freshly initialized module values are kept (and remain visible to the + subsequent strict ``load_state_dict``). + +The retargeted sharded state dict is then handed to the regular +``dist_checkpointing.load`` machinery, which reads the GPT checkpoint +directly and reshards across any TP/PP/EP/ETP layout change on the way. + +The same retargeting also applies to the distributed optimizer's sharded +state dict. In the model-space checkpoint formats (``fully_reshardable`` / +``fully_sharded_model_space``) every optimizer-state ``ShardedTensor`` is built +by copying the corresponding model param's metadata and prefixing its ``key`` +with ``optimizer.state..`` (see +``DistributedOptimizer.sharded_param_state_*``). Those entries therefore carry +the same ``decoder.layers..`` keys and sharding as the model tensors, so +:func:`retarget_sharded_state_dict_to_gpt_checkpoint` rewrites them onto the GPT +checkpoint identically -- optimizer moments and fp32 master params for +attention/MLP layers load from the GPT run, while fresh layers (e.g. Mamba) +keep their freshly initialized optimizer state via ``LocalNonpersistentObject``. +""" + +import re +from dataclasses import dataclass +from typing import Any, Iterable, Mapping + +from megatron.core.dist_checkpointing.dict_utils import dict_list_map_inplace +from megatron.core.dist_checkpointing.mapping import ( + LocalNonpersistentObject, + ShardedBase, + ShardedObject, + ShardedStateDict, + ShardedTensor, + ShardedTensorFactory, +) +from megatron.core.models.hybrid.hybrid_layer_allocation import ( + Symbols, + get_layer_maps_from_layer_type_list, + parse_hybrid_pattern, +) + +# Hybrid layer symbols that have a GPT-side source of weights ('*', '-', 'E') +# or that are explicitly initialized from scratch ('M'). Attention layers map +# onto GPT ``self_attention`` sub-modules; dense and MoE MLP layers map onto +# GPT ``mlp`` sub-modules (both models keep MoE tensors under ``mlp.*``). +# GDN ('G') and DS-attention ('D') use different weight layouts than GPT +# attention and are rejected rather than silently mistranslated. +_GPT_SOURCED_SYMBOLS = (Symbols.ATTENTION, Symbols.MLP, Symbols.MOE) +_FRESH_INIT_SYMBOLS = (Symbols.MAMBA,) + +_DECODER_LAYER_KEY_RE = re.compile(r'decoder\.layers\.(\d+)\.') + +_GPT_FINAL_NORM_KEY_MAP = {'decoder.final_norm.': 'decoder.final_layernorm.'} +_GPT_OMITTED_LOCAL_KEYS = ('output_layer._extra_state',) + + +@dataclass(frozen=True) +class GPTCompatLayerMaps: + """Correspondence between hybrid layer indices and GPT layer indices. + + Attributes: + attention_to_gpt: hybrid global layer index of the i-th attention + position -> GPT layer index i. + mlp_to_gpt: hybrid global layer index of the i-th MLP-bearing + position ('-' or 'E') -> GPT layer index i. + fresh_init: hybrid global layer indices with no GPT counterpart; + their modules keep the run's fresh initialization. + num_gpt_layers: number of layers the source GPT checkpoint must have. + """ + + attention_to_gpt: Mapping[int, int] + mlp_to_gpt: Mapping[int, int] + fresh_init: frozenset + num_gpt_layers: int + + +def gpt_compatible_layer_maps(hybrid_layer_pattern: str) -> GPTCompatLayerMaps: + """Derive hybrid->GPT layer index maps from a hybrid layer pattern. + + Args: + hybrid_layer_pattern: the run's unified hybrid layer pattern + (pipeline '|' separators allowed). + + Returns: + GPTCompatLayerMaps for retargeting a sharded state dict. + + Raises: + ValueError: if the pattern cannot be paired one-to-one with a GPT + checkpoint layout (MTP present, non-translatable symbols, + mixed dense/MoE positions, or unbalanced '*' vs MLP counts). + """ + parsed = parse_hybrid_pattern(hybrid_layer_pattern) + if parsed.mtp_num_depths > 0: + raise ValueError( + f"Hybrid layer pattern {hybrid_layer_pattern!r} contains MTP layers " + f"('/{parsed.mtp_pattern}'), which have no source weights in a GPT " + f"checkpoint. Remove the MTP part of the pattern to load a GPT checkpoint." + ) + main_pattern = (parsed.main_pattern or '').replace(Symbols.PIPE, '') + if not main_pattern: + raise ValueError("Hybrid layer pattern is empty; set --hybrid-layer-pattern.") + + layer_type_list = list(main_pattern) + translatable = set(_GPT_SOURCED_SYMBOLS) | set(_FRESH_INIT_SYMBOLS) + unknown = sorted(set(layer_type_list) - translatable) + if unknown: + raise ValueError( + f"Hybrid layer pattern {hybrid_layer_pattern!r} contains layer types " + f"{unknown} that cannot be translated from a GPT checkpoint. " + f"Supported: {sorted(translatable)} ('M' layers keep their fresh " + f"initialization)." + ) + + layer_maps = get_layer_maps_from_layer_type_list(layer_type_list) + dense_map = layer_maps[Symbols.MLP] + moe_map = layer_maps[Symbols.MOE] + if dense_map and moe_map: + raise ValueError( + f"Hybrid layer pattern {hybrid_layer_pattern!r} mixes dense ('-') and " + f"MoE ('E') MLP positions. GPT checkpoints have one MLP kind on every " + f"layer, so the pattern must use only one of '-' or 'E'." + ) + mlp_map = moe_map if moe_map else dense_map + attention_map = layer_maps[Symbols.ATTENTION] + + if len(attention_map) != len(mlp_map) or not attention_map: + raise ValueError( + f"Hybrid layer pattern {hybrid_layer_pattern!r} has " + f"{len(attention_map)} attention ('*') and {len(mlp_map)} MLP ('-'/'E') " + f"positions. Each GPT layer provides exactly one attention and one MLP " + f"sub-module, so the pattern needs an equal, nonzero number of each." + ) + + return GPTCompatLayerMaps( + attention_to_gpt=dict(attention_map), + mlp_to_gpt=dict(mlp_map), + fresh_init=frozenset(layer_maps[Symbols.MAMBA]), + num_gpt_layers=len(attention_map), + ) + + +def _prepend_gpt_layer_axis(entry, gpt_layer_idx: int, num_gpt_layers: int): + """Add the GPT layer index as the leading sharding axis of an entry. + + Mirrors what ``TransformerBlock.sharded_state_dict`` does for homogeneous + layers by passing ``sharded_offsets=[(0, layer_idx, num_layers)]`` down to + ``make_sharded_tensors_for_checkpoint``: + + * ShardedTensor: one more prepended axis of size ``num_gpt_layers`` + at position 0, this shard sitting at ``gpt_layer_idx``; + * ShardedObject: ``(1,)/(0,)`` placeholder offsets (from + ``_get_extra_state_offsets`` with no offsets) are replaced by the layer + axis, otherwise the layer axis is prepended (e.g. before an expert axis); + * ShardedTensorFactory: the built sub-entries get the same treatment. + """ + if isinstance(entry, ShardedTensor): + entry.global_shape = (num_gpt_layers, *entry.global_shape) + entry.global_offset = (gpt_layer_idx, *entry.global_offset) + entry.axis_fragmentations = (num_gpt_layers, *entry.axis_fragmentations) + entry.prepend_axis_num += 1 + elif isinstance(entry, ShardedObject): + if entry.global_shape == (1,) and entry.global_offset == (0,): + entry.global_shape = (num_gpt_layers,) + entry.global_offset = (gpt_layer_idx,) + else: + entry.global_shape = (num_gpt_layers, *entry.global_shape) + entry.global_offset = (gpt_layer_idx, *entry.global_offset) + elif isinstance(entry, ShardedTensorFactory): + inner_build_fn = entry.build_fn + + def _build_with_gpt_layer_axis(key, data, replica_id, flattened_range): + built = inner_build_fn(key, data, replica_id, flattened_range) + dict_list_map_inplace( + lambda sub: _prepend_gpt_layer_axis(sub, gpt_layer_idx, num_gpt_layers), built + ) + return built + + entry.build_fn = _build_with_gpt_layer_axis + return entry + + +def retarget_sharded_state_dict_to_gpt_checkpoint( + sharded_state_dict: ShardedStateDict, layer_maps: GPTCompatLayerMaps +) -> None: + """Point a hybrid model's sharded state dict at a GPT checkpoint, in place. + + Only the storage lookup metadata (``key`` and sharding axes) of each + ``ShardedBase`` entry is rewritten into the GPT checkpoint's homogeneous + layer format; the nested state dict structure (used by the subsequent + ``load_state_dict``) keeps the hybrid model's own names. Entries of layers + with no GPT counterpart are replaced by ``LocalNonpersistentObject`` so the + loaded state dict returns their current (freshly initialized) values. + + The same routine handles the distributed optimizer's sharded state dict: its + per-parameter entries embed the model key (``optimizer.state..decoder. + layers....``) and mirror the model param's sharding, so they retarget the + same way, and fresh-layer optimizer state is likewise kept local. + + Args: + sharded_state_dict: one model chunk's sharded state dict (as produced by + ``model.sharded_state_dict()``) or the matching optimizer sharded + state dict. + layer_maps: maps from :func:`gpt_compatible_layer_maps` derived from + the same pattern the model was built with. + """ + + def _retarget(entry): + if not isinstance(entry, ShardedBase): + return entry + + if entry.key.endswith(_GPT_OMITTED_LOCAL_KEYS): + return LocalNonpersistentObject(entry.data) + + layer_match = _DECODER_LAYER_KEY_RE.search(entry.key) + if layer_match is not None: + hybrid_idx = int(layer_match.group(1)) + if hybrid_idx in layer_maps.fresh_init: + return LocalNonpersistentObject(entry.data) + gpt_idx = layer_maps.attention_to_gpt.get(hybrid_idx) + if gpt_idx is None: + gpt_idx = layer_maps.mlp_to_gpt.get(hybrid_idx) + if gpt_idx is None: + raise ValueError( + f"Sharded state dict entry {entry.key!r} refers to hybrid layer " + f"{hybrid_idx}, which is not part of the hybrid layer pattern " + f"used to derive the GPT layer maps. The pattern and the " + f"instantiated model do not match." + ) + # GPT checkpoints use the homogeneous layer format: no layer index + # in the key, the layer is a sharding axis instead. + entry.key = ( + f'{entry.key[:layer_match.start()]}decoder.layers.' + f'{entry.key[layer_match.end():]}' + ) + return _prepend_gpt_layer_axis(entry, gpt_idx, layer_maps.num_gpt_layers) + + for hybrid_prefix, gpt_prefix in _GPT_FINAL_NORM_KEY_MAP.items(): + pos = entry.key.find(hybrid_prefix) + if pos != -1: + entry.key = f'{entry.key[:pos]}{gpt_prefix}{entry.key[pos + len(hybrid_prefix):]}' + break + return entry + + dict_list_map_inplace(_retarget, sharded_state_dict) + + +def _retarget_explicit_key_to_gpt_checkpoint( + key: Any, layer_maps: GPTCompatLayerMaps, checkpoint_keys: Iterable[str] | None = None +) -> Any | None: + """Translate one explicit HybridModel state-dict key to its GPT key. + + ``fsdp_dtensor`` checkpoints store explicit parameter names rather than + homogeneous-layer ``ShardedTensor`` metadata. Returning ``None`` omits a + fresh-only or GPT-omitted entry from the DCP load plan while leaving its + existing HybridModel value untouched. + """ + if not isinstance(key, str): + return key + if key.endswith(_GPT_OMITTED_LOCAL_KEYS): + return None + + layer_match = _DECODER_LAYER_KEY_RE.search(key) + if layer_match is not None: + hybrid_idx = int(layer_match.group(1)) + if hybrid_idx in layer_maps.fresh_init: + return None + gpt_idx = layer_maps.attention_to_gpt.get(hybrid_idx) + if gpt_idx is None: + gpt_idx = layer_maps.mlp_to_gpt.get(hybrid_idx) + if gpt_idx is None: + raise ValueError( + f"FSDP state dict entry {key!r} refers to hybrid layer {hybrid_idx}, " + "which is not part of the hybrid layer pattern used to derive the " + "GPT layer maps." + ) + key = f'{key[:layer_match.start()]}decoder.layers.{gpt_idx}.' f'{key[layer_match.end():]}' + + for hybrid_prefix, gpt_prefix in _GPT_FINAL_NORM_KEY_MAP.items(): + pos = key.find(hybrid_prefix) + if pos != -1: + key = f'{key[:pos]}{gpt_prefix}{key[pos + len(hybrid_prefix):]}' + break + + # FSDP optimizer parameter names include the wrapper hierarchy. GPTModel + # and HybridModel can have different Float16/FSDP wrapper depths, so use + # checkpoint metadata to recover the exact source-side ``module.`` prefix. + if checkpoint_keys is not None: + bare_key = re.sub(r'^(?:module\.)+', '', key) + key_pattern = re.compile(rf'(? dict[Any, Any]: + """Return an ``fsdp_dtensor`` model or optimizer state dict under GPT keys. + + FSDP model state is a flat parameter-name mapping. Distributed-optimizer + state can contain nested ``state`` and ``param_to_group_meta`` mappings + (and chained-optimizer integer keys), so the translation recursively + rewrites every parameter-name key while preserving the DTensor leaves. + """ + + checkpoint_key_set = set(checkpoint_keys) if checkpoint_keys is not None else None + + def _retarget(value, path): + if isinstance(value, Mapping): + translated = {} + for key, child in value.items(): + translated_key = _retarget_explicit_key_to_gpt_checkpoint( + key, layer_maps, checkpoint_key_set + ) + if translated_key is not None: + child_path = f'{path}.{translated_key}' if path else str(translated_key) + translated_child = _retarget(child, child_path) + if ( + checkpoint_key_set is None + or isinstance(child, (Mapping, list, tuple)) + or child_path in checkpoint_key_set + ): + translated[translated_key] = translated_child + return translated + if isinstance(value, list): + return [_retarget(child, f'{path}.{idx}') for idx, child in enumerate(value)] + if isinstance(value, tuple): + return tuple(_retarget(child, f'{path}.{idx}') for idx, child in enumerate(value)) + return value + + return _retarget(state_dict, checkpoint_prefix) diff --git a/megatron/core/dist_checkpointing/serialization.py b/megatron/core/dist_checkpointing/serialization.py index dd85fe178cd..cc08aa26dbf 100644 --- a/megatron/core/dist_checkpointing/serialization.py +++ b/megatron/core/dist_checkpointing/serialization.py @@ -1,4 +1,4 @@ -# Copyright (c) 2022-2023, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. """Entrypoints for saving and loading the distributed checkpoints. @@ -16,7 +16,7 @@ import torch -from megatron.core.msc_utils import MultiStorageClientFeature +from megatron.core.msc_utils import maybe_msc from . import ShardedTensor from .core import CheckpointingConfig, save_config @@ -49,9 +49,16 @@ logger = logging.getLogger(__name__) -# monkeypatch needed for ModelOpt -# will be removed once MLM updated to newer ModelOpt -get_default_load_sharded_strategy = TorchDistLoadShardedStrategy + +def get_default_load_sharded_strategy(checkpoint_dir: str | Path | None = None): + """Create the default torch distributed load strategy.""" + return TorchDistLoadShardedStrategy(checkpoint_name=checkpoint_dir) + + +def get_default_save_sharded_strategy(backend: str = "torch_dist"): + """Create the default torch distributed save strategy.""" + return TorchDistSaveShardedStrategy(backend=backend) + # flat state dict with sharded objects without any data CkptShardedMetadata = Dict[str, Union[ShardedTensor, ShardedObject]] @@ -66,6 +73,7 @@ def load( validate_access_integrity: bool = True, strict: Union[str, StrictHandling] = StrictHandling.ASSUME_OK_UNEXPECTED, verify_integrity: bool = False, + process_group: Optional[torch.distributed.ProcessGroup] = None, ) -> Union[StateDict, Tuple[StateDict, Set[str], Set[str]]]: """Loading entrypoint. @@ -101,6 +109,8 @@ def load( and compares against the SHA-256 manifest. Raises `CheckpointingException` on any mismatch. Requires that the checkpoint was previously saved with `verify_integrity=True`. + process_group (ProcessGroup, optional): ranks that collectively describe + one complete sharded state dict. Defaults to the global process group. Returns: StateDict or Tuple[StateDict, Set[str], Set[str]]: in most cases only @@ -122,7 +132,15 @@ def load( # params with a high-precision state dict; # 2. When using delayed scaling, this loading process writes an extra value into the global # amax_history buffer of Transformer Engine, which is undesirable. - force_all_tensors_to_non_fp8(sharded_state_dict) + # + # When the sharded strategy supports per-tensor streaming dequantize + # (``stream_ckpt_dequant``), both concerns are handled inside the + # LoadPlanner on a per-tensor basis, which avoids peaking GPU memory + # with N simultaneous high-precision scratch tensors before the load + # begins. Covers FP8/MXFP8/blockwise-FP8/NVFP4 via the common + # ``QuantizedTensor`` base class. + if not getattr(sharded_strategy, "stream_ckpt_dequant", False): + force_all_tensors_to_non_fp8(sharded_state_dict) sharded_state_dict, nonpersistent_state_dict, sh_ten_factories = load_preprocess( sharded_state_dict @@ -148,7 +166,9 @@ def load( k: v for k, v in ckpt_sharded_metadata.items() if v.key != 'common_state' } if validate_access_integrity or StrictHandling.requires_global_app_metadata(strict): - local_metadata, global_metadata = determine_global_metadata(sharded_state_dict) + local_metadata, global_metadata = determine_global_metadata( + sharded_state_dict, process_group=process_group + ) sharded_state_dict, missing_keys, unexpected_keys = validate_integrity_and_strict_load( sharded_state_dict, @@ -180,10 +200,7 @@ def load( def _legacy_common_state_exists(checkpoint_dir: str) -> bool: """Check whether the checkpoint stores common data in a legacy common.pt file.""" path = os.path.join(checkpoint_dir, COMMON_STATE_FNAME) - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - return msc.Path(path).exists() - return os.path.exists(path) + return maybe_msc.Path(path).exists() def load_common_state_dict(checkpoint_dir: Union[str, Path]) -> StateDict: @@ -326,7 +343,7 @@ def load_content_metadata( def remove_sharded_tensors(checkpoint_dir: str, key_prefix: str): """determine the appropriate sharding strategy and delegate removal to the sharded strategy""" verify_checkpoint(checkpoint_dir) - TorchDistSaveShardedStrategy.remove_sharded_tensors(checkpoint_dir, key_prefix) + TorchDistLoadShardedStrategy().remove_sharded_tensors(checkpoint_dir, key_prefix) def save( @@ -398,11 +415,7 @@ def save( from .strategies.fully_parallel import FullyParallelSaveStrategyWrapper if torch.distributed.get_rank() == 0: - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - checkpoint_dir_path = msc.Path(str(checkpoint_dir)) - else: - checkpoint_dir_path = Path(checkpoint_dir) + checkpoint_dir_path = maybe_msc.Path(str(checkpoint_dir)) if next(checkpoint_dir_path.iterdir(), None) is not None: # Don't throw exception here since this could cause a cascade of failures diff --git a/megatron/core/dist_checkpointing/strategies/common.py b/megatron/core/dist_checkpointing/strategies/common.py index 1ec3d829275..3e9685e2b9c 100644 --- a/megatron/core/dist_checkpointing/strategies/common.py +++ b/megatron/core/dist_checkpointing/strategies/common.py @@ -4,12 +4,11 @@ import logging import os -from pathlib import Path import torch from megatron.core.dist_checkpointing.mapping import StateDict -from megatron.core.msc_utils import MultiStorageClientFeature +from megatron.core.msc_utils import maybe_msc from ..mapping import CheckpointingException @@ -27,11 +26,7 @@ def save_common(common_state_dict: StateDict, checkpoint_dir: str): if torch.distributed.get_rank() == 0: path = os.path.join(checkpoint_dir, COMMON_STATE_FNAME) - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - msc.torch.save(common_state_dict, path) - else: - torch.save(common_state_dict, path) + maybe_msc.torch.save(common_state_dict, path) def load_common(checkpoint_dir: str): @@ -50,17 +45,9 @@ def load_common(checkpoint_dir: str): load_path = os.path.join(checkpoint_dir, COMMON_STATE_FNAME) try: - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - return msc.torch.load(load_path, map_location='cpu') - else: - return torch.load(load_path, map_location='cpu') + return maybe_msc.torch.load(load_path, map_location='cpu') except FileNotFoundError as e: err_msg = f'Common file {load_path} does not exist' - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - ckpt_files = [f.name for f in msc.Path(checkpoint_dir).iterdir()] - else: - ckpt_files = [f.name for f in Path(checkpoint_dir).iterdir()] + ckpt_files = [f.name for f in maybe_msc.Path(checkpoint_dir).iterdir()] logger.debug(f'{err_msg}. Checkpoint directory content: {ckpt_files}') raise CheckpointingException(err_msg) from e diff --git a/megatron/core/dist_checkpointing/strategies/fully_parallel.py b/megatron/core/dist_checkpointing/strategies/fully_parallel.py index db3c8ee6cae..c201224efe9 100644 --- a/megatron/core/dist_checkpointing/strategies/fully_parallel.py +++ b/megatron/core/dist_checkpointing/strategies/fully_parallel.py @@ -184,6 +184,13 @@ def __init__( self.cached_distribution: Optional[ShardDistribution] = None self.cached_global_metadata: Optional[Metadata] = None + @property + def stream_ckpt_dequant(self) -> bool: + """Forward the streaming dequantize flag from the wrapped strategy so that + ``serialization.load`` can skip the upfront ``force_all_tensors_to_non_fp8`` pass + when streaming is enabled.""" + return getattr(self.base_strategy, "stream_ckpt_dequant", False) + @debug_time("FullyParallelLoadStrategyWrapper.load", logger) def load( self, diff --git a/megatron/core/dist_checkpointing/strategies/torch.py b/megatron/core/dist_checkpointing/strategies/torch.py index 8d65e299304..d8487c4bb43 100644 --- a/megatron/core/dist_checkpointing/strategies/torch.py +++ b/megatron/core/dist_checkpointing/strategies/torch.py @@ -1,4 +1,4 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. """ Strategies using PyTorch distributed.checkpoint as an underlying format. """ import inspect @@ -349,9 +349,20 @@ def _unwrap_pyt_sharded_tensor( ret_tensors = [] for sh in sh_ten.local_shards(): ten = sh.tensor - for _ in range(mcore_sh_ten.prepend_axis_num): - assert ten.size(0) == 1 - ten = ten[0] # NOTE: ten.squeeze(0) uses more memory for FP8 tensors + if mcore_sh_ten.prepend_axis_num > 0: + # NOTE: use ``view`` to strip the prepended singleton axes. Indexing + # (``ten[0]``) and ``squeeze`` both dispatch through + # ``aten.select.int`` / ``aten.squeeze`` which are not implemented + # by ``MXFP8Tensor`` (nor by the blockwise tensor class) — they + # fall back to ``QuantizedTensor.__torch_dispatch__`` which + # dequantizes the entire tensor to BF16 before applying the op. + # For large FP8/MXFP8 checkpoints this materializes a full-size + # BF16 copy per shard and OOMs right after a successful streaming + # load. ``view`` is handled natively by all TE quantized tensor + # classes without any dequantize. + for i in range(mcore_sh_ten.prepend_axis_num): + assert ten.size(i) == 1 + ten = ten.view(ten.shape[mcore_sh_ten.prepend_axis_num :]) ret_tensors.append(ten) return ret_tensors @@ -491,12 +502,19 @@ def __init__( *args, shapes_validation_sharded_tensors: Iterable[ShardedTensor] = (), allow_shape_mismatch_sharded_tensors: Optional[Dict[str, ShardedTensor]] = None, + stream_ckpt_dequant: bool = True, **kwargs, ) -> None: super().__init__(*args, **kwargs) self.shapes_validation_sharded_tensors = shapes_validation_sharded_tensors self.allow_shape_mismatch_sharded_tensors = allow_shape_mismatch_sharded_tensors - self._intermediate_read_item_and_target: Optional[Tuple[ReadItem, torch.Tensor]] = None + self.stream_ckpt_dequant = stream_ckpt_dequant + # Maps id(read_item) -> (read_item, target_tensor, amax_snapshot_or_None, kind) + # kind is "stream" for the streaming per-tensor dequant path and "noncontig" + # for the existing contiguity-fix path. + self._intermediate_read_items: Dict[ + int, Tuple[ReadItem, torch.Tensor, Optional[torch.Tensor], str] + ] = {} def _validate_global_shapes(self, metadata, sharded_tensors): for sh_ten in sharded_tensors: @@ -550,37 +568,92 @@ def create_local_plan(self) -> LoadPlan: return local_plan def resolve_tensor(self, read_item: ReadItem): - """Override to add FP8 support. - - Narrowing the Float8Tensor can create incontiguous tensors and there are - no `copy` kernels for such cases. This method creates a contiguous FP8 - tensors so that the subsequent `copy_` in FileSystemReader succeeds. - Note that this requires tracking the original tensor - (as `self._intermediate_read_item_and_target` attribute) - and restoring it in `commit_tensor` method. + """Override to add quantized-tensor support. + + Two paths are handled here: + + 1. Streaming per-tensor dequantize (when ``stream_ckpt_dequant`` is True + and the destination is a TE ``QuantizedTensor`` — covers Float8, + MXFP8, blockwise FP8, and NVFP4 via the common base class). We + allocate a per-tensor high-precision scratch buffer, return it as + the load destination, and quantize-copy it back into the original + tensor in ``commit_tensor``. This replaces the upfront bulk + dequantize done by ``force_all_tensors_to_non_fp8`` and keeps at + most one scratch tensor live at a time. + + 2. Non-contiguous Float8 fix: narrowing a Float8Tensor can produce a + non-contiguous view for which no ``copy_`` kernel exists. We fall + back to a contiguous Float8 clone and copy it back in + ``commit_tensor``. + + Both cases stash state in ``self._intermediate_read_items``, keyed by + ``id(read_item)``, so ``commit_tensor`` can undo them. """ target_tensor = super().resolve_tensor(read_item) + + # Lazy import to avoid circular imports (fp8_utils pulls in core.tensor_parallel). + from ...fp8_utils import is_float8tensor as _is_quantized_tensor + + if ( + self.stream_ckpt_dequant + and HAVE_TE + and _is_quantized_tensor(target_tensor) + and target_tensor.is_cuda + ): + # Snapshot amax for delayed-scaling quantizers so the subsequent + # BF16->FP8 quantize-copy does not pollute amax_history. For + # current-scaling / MXFP8 / blockwise / NVFP4 quantizers, amax is + # None or absent on the quantizer and the snapshot is a no-op. + amax_snapshot: Optional[torch.Tensor] = None + quantizer = getattr(target_tensor, "_quantizer", None) + amax = getattr(quantizer, "amax", None) if quantizer is not None else None + if isinstance(amax, torch.Tensor): + amax_snapshot = amax.detach().clone() + + scratch = torch.empty( + target_tensor.shape, dtype=target_tensor.dtype, device=target_tensor.device + ) + self._intermediate_read_items[id(read_item)] = ( + read_item, + target_tensor, + amax_snapshot, + "stream", + ) + return scratch + if ( not target_tensor.is_contiguous() and HAVE_TE and isinstance(target_tensor, Float8Tensor) ): - self._intermediate_read_item_and_target = (read_item, target_tensor) + self._intermediate_read_items[id(read_item)] = ( + read_item, + target_tensor, + None, + "noncontig", + ) target_tensor = Float8Tensor.make_like( target_tensor, data=target_tensor._data.contiguous() ) return target_tensor def commit_tensor(self, read_item: ReadItem, tensor: torch.Tensor) -> None: - """Restores the original FP8 tensor saved in `resolve_tensor`.""" - if self._intermediate_read_item_and_target is not None: - interm_read_item, target_tensor = self._intermediate_read_item_and_target - assert ( - interm_read_item is read_item - ), '`commit_tensor` method should be called right after `resolve_tensor`' + """Undo the detours stashed in ``resolve_tensor``. + + - Streaming case: copy the high-precision scratch back into the + original quantized tensor (quantize-on-copy), then restore the + pre-load ``amax`` for delayed-scaling quantizers. + - Non-contiguous case: copy the contiguous clone back into the + original narrowed Float8Tensor view. + """ + entry = self._intermediate_read_items.pop(id(read_item), None) + if entry is not None: + _, target_tensor, amax_snapshot, kind = entry target_tensor.copy_(tensor) + if kind == "stream" and amax_snapshot is not None: + # quantizer was non-None when we took the snapshot + target_tensor._quantizer.amax.copy_(amax_snapshot) tensor = target_tensor - self._intermediate_read_item_and_target = None return super().commit_tensor(read_item, tensor) @@ -851,9 +924,20 @@ def _get_filesystem_reader( class TorchDistLoadShardedStrategy: """Basic load strategy for the PyT Distributed format.""" - def __init__(self, cache_metadata: bool = False, checkpoint_name: str = None): + def __init__( + self, + cache_metadata: bool = False, + stream_ckpt_dequant: bool = True, + checkpoint_name: str = None, + ): self.cached_global_metadata: Optional[Metadata] = None self.cache_metadata = cache_metadata + # When True, quantized destinations (FP8/MXFP8/blockwise FP8/NVFP4) are + # dequantized per-tensor inside the LoadPlanner rather than all at once + # before the load starts. This trades a small planner overhead for a + # large reduction in peak GPU memory during load. See + # serialization.load() and MCoreLoadPlanner.resolve_tensor for details. + self.stream_ckpt_dequant = stream_ckpt_dequant self.checkpoint_name = checkpoint_name def load( @@ -898,6 +982,7 @@ def load( planner=MCoreLoadPlanner( shapes_validation_sharded_tensors=flexible_shape_sharded_tensors, allow_shape_mismatch_sharded_tensors=allow_shape_mismatch_sharded_tensors, + stream_ckpt_dequant=self.stream_ckpt_dequant, flatten_state_dict=False, flatten_sharded_tensors=False, ), @@ -1024,16 +1109,16 @@ def remove_sharded_tensors(self, checkpoint_dir: str, key_prefix: str): except AttributeError: os.sync() ## move the old metadata - fs_writer.fs.rename(fs_writer.metadata_path, old_path) + fs_writer.fs.rename(metadata_filename, old_path) try: ## rename the new metadata - fs_writer.fs.rename(tmp_path, fs_writer.metadata_path) + fs_writer.fs.rename(tmp_path, metadata_filename) ## finally, remove the files we want to drop for f in files_to_remove: - fs_writer.fs.rm_file(checkpoint_dir / f) + fs_writer.fs.rm_file(Path(checkpoint_dir) / f) except Exception as e: - fs_writer.fs.rename(old_path, fs_writer.metadata_path) + fs_writer.fs.rename(old_path, metadata_filename) raise e else: fs_writer.fs.rm_file(old_path) diff --git a/megatron/core/dist_checkpointing/validation.py b/megatron/core/dist_checkpointing/validation.py index b0cbae618a7..6095ea95d8f 100644 --- a/megatron/core/dist_checkpointing/validation.py +++ b/megatron/core/dist_checkpointing/validation.py @@ -6,7 +6,6 @@ import os from collections import Counter, defaultdict from enum import Enum -from pathlib import Path from typing import TYPE_CHECKING, Dict, List, Optional, Set, Tuple, Union import numpy as np @@ -25,7 +24,7 @@ ShardedStateDict, is_main_replica, ) -from megatron.core.msc_utils import MultiStorageClientFeature +from megatron.core.msc_utils import maybe_msc if TYPE_CHECKING: from megatron.core.dist_checkpointing.serialization import CkptShardedMetadata @@ -207,12 +206,7 @@ def verify_checkpoint(checkpoint_dir: str): Args: checkpoint_dir (str): checkpoint directory """ - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - isdir = msc.os.path.isdir(str(checkpoint_dir), strict=False) - else: - isdir = os.path.isdir(checkpoint_dir) - if not isdir: + if not maybe_msc.path_isdir(str(checkpoint_dir), strict=False): raise CheckpointingException(f'Checkpoint directory {checkpoint_dir} does not exist') if not check_is_distributed_checkpoint(checkpoint_dir): @@ -483,18 +477,21 @@ def _validate_objects_for_key(sharded_objects: List[ShardedObject]) -> List[Chec def determine_global_metadata( sharded_state_dict: ShardedStateDict, + process_group: Optional[torch.distributed.ProcessGroup] = None, ) -> Tuple[_LocalMetadata, _GlobalMetadata]: """Exchanges local metadata with `all_gather_object` to determine global metadata. Args: sharded_state_dict (ShardedStateDict): local sharded state dict + process_group (ProcessGroup, optional): ranks whose metadata forms one + complete checkpoint view. Defaults to the global process group. Returns: Tuple[_LocalMetadata, _GlobalMetadata]: local and global ShardedBase objects with stripped data """ local_metadata = [ten.without_data() for ten in nested_values(sharded_state_dict)] - global_metadata = [None] * torch.distributed.get_world_size() - torch.distributed.all_gather_object(global_metadata, local_metadata) + global_metadata = [None] * torch.distributed.get_world_size(group=process_group) + torch.distributed.all_gather_object(global_metadata, local_metadata, group=process_group) return local_metadata, global_metadata # type: ignore[return-value] @@ -506,15 +503,9 @@ def _compute_file_hash(file_path: str) -> str: Lowercase hex-encoded SHA-256 digest string. """ h = hashlib.sha256() - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - with msc.open(file_path, 'rb') as f: - for chunk in iter(lambda: f.read(_READ_CHUNK_SIZE), b''): - h.update(chunk) - else: - with open(file_path, 'rb') as f: - for chunk in iter(lambda: f.read(_READ_CHUNK_SIZE), b''): - h.update(chunk) + with maybe_msc.open(file_path, 'rb') as f: + for chunk in iter(lambda: f.read(_READ_CHUNK_SIZE), b''): + h.update(chunk) return h.hexdigest() @@ -528,28 +519,16 @@ def save_integrity_manifest(checkpoint_dir: str) -> None: """ manifest: Dict[str, str] = {} - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - ckpt_path = msc.Path(checkpoint_dir) - for entry in sorted(ckpt_path.iterdir(), key=lambda p: str(p)): - if entry.name != INTEGRITY_FNAME: - manifest[entry.name] = _compute_file_hash(str(entry)) - else: - ckpt_path = Path(checkpoint_dir) - for entry in sorted(ckpt_path.iterdir()): - if entry.is_file() and entry.name != INTEGRITY_FNAME: - manifest[entry.name] = _compute_file_hash(str(entry)) + ckpt_path = maybe_msc.Path(checkpoint_dir) + for entry in sorted(ckpt_path.iterdir(), key=lambda p: str(p)): + if entry.is_file() and entry.name != INTEGRITY_FNAME: + manifest[entry.name] = _compute_file_hash(str(entry)) integrity_path = os.path.join(checkpoint_dir, INTEGRITY_FNAME) payload = {'algorithm': _HASH_ALGORITHM, 'files': manifest} - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - with msc.open(integrity_path, 'w') as f: - json.dump(payload, f, indent=2) - else: - with open(integrity_path, 'w') as f: - json.dump(payload, f, indent=2) + with maybe_msc.open(integrity_path, 'w') as f: + json.dump(payload, f, indent=2) logger.info("Saved integrity manifest with %d file(s) to %s", len(manifest), integrity_path) @@ -567,25 +546,14 @@ def _verify_integrity_manifest_impl(checkpoint_dir: str) -> None: """ integrity_path = os.path.join(checkpoint_dir, INTEGRITY_FNAME) - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - if not msc.os.path.exists(integrity_path): - raise CheckpointingException( - f'Integrity manifest not found at {integrity_path}. ' - 'The checkpoint must be saved with integrity verification enabled ' - '(save_integrity=True) before it can be verified on load.' - ) - with msc.open(integrity_path) as f: - manifest_data = json.load(f) - else: - if not os.path.exists(integrity_path): - raise CheckpointingException( - f'Integrity manifest not found at {integrity_path}. ' - 'The checkpoint must be saved with integrity verification enabled ' - '(save_integrity=True) before it can be verified on load.' - ) - with open(integrity_path) as f: - manifest_data = json.load(f) + if not maybe_msc.os.path.exists(integrity_path): + raise CheckpointingException( + f'Integrity manifest not found at {integrity_path}. ' + 'The checkpoint must be saved with integrity verification enabled ' + '(save_integrity=True) before it can be verified on load.' + ) + with maybe_msc.open(integrity_path) as f: + manifest_data = json.load(f) algorithm = manifest_data.get('algorithm', _HASH_ALGORITHM) if algorithm != _HASH_ALGORITHM: diff --git a/megatron/core/distributed/distributed_data_parallel.py b/megatron/core/distributed/distributed_data_parallel.py index c4957250930..90d130dedd8 100644 --- a/megatron/core/distributed/distributed_data_parallel.py +++ b/megatron/core/distributed/distributed_data_parallel.py @@ -174,6 +174,15 @@ def __init__( self.full_param_layout = full_param_layout + # GTP_remat needs average_in_collective=False: the per-bucket collective runs over the + # replicate group, so NCCL AVG would miss the 1/gtp_remat factor. arguments.py + # guards the training path; this assert covers direct megatron-core users. + gtp_active = ProcessGroupCollection.is_gtp_remat_active(process_group_dict) + assert not (gtp_active and self.ddp_config.average_in_collective), ( + "GTP requires average_in_collective=False (the default); averaged collectives reduce " + "over the GTP-excluded group and would miss the 1/gtp_remat gradient scaling factor." + ) + # Compute gradient scaling factors. if config.calculate_per_token_loss: assert ( @@ -372,6 +381,13 @@ def unmap_weight_tensor(m): self._make_backward_post_hook(param) ) break + elif getattr(param, 'is_gtp_weight_remat', False) and hasattr( + param, 'register_grad_accum_hook' + ): + # GTP_remat defers the main_grad add to a later backward node, so drive the + # post-hook from its manual call (_handle_megatron_grad_accum) rather than + # autograd's AccumulateGrad, which would fire grad-ready on stale main_grad. + param.register_grad_accum_hook(None, self._make_backward_post_hook(param)) else: # Expand so we get access to grad_fn. param_tmp = param.expand_as(param) @@ -472,9 +488,13 @@ def hook(*unused): assert param.requires_grad cudagraph_wgrad_ready_event = getattr(param, '_cudagraph_wgrad_ready_event', None) if self.ddp_config.overlap_grad_reduce and cudagraph_wgrad_ready_event is None: - assert ( - param.grad is not None - ), 'param.grad being None is not safe when overlap_grad_reduce is True' + # GTP_remat keeps its real wgrad in main_grad (via finalize); param.grad here is + # throwaway (None or a dummy), so skip this assert and rely on + # grad_added_to_main_grad below. + if not getattr(param, 'is_gtp_weight_remat', False): + assert ( + param.grad is not None + ), 'param.grad being None is not safe when overlap_grad_reduce is True' if param.grad is not None and ( not param.grad_added_to_main_grad or getattr(param, 'zero_out_wgrad', False) ): diff --git a/megatron/core/distributed/finalize_model_grads.py b/megatron/core/distributed/finalize_model_grads.py index e494bbde2b6..15f0355bd5b 100644 --- a/megatron/core/distributed/finalize_model_grads.py +++ b/megatron/core/distributed/finalize_model_grads.py @@ -330,6 +330,9 @@ def reset_model_temporary_tensors(config: TransformerConfig, model: List[torch.n or "global_aux_loss" in config.moe_router_load_balancing_type ) and hasattr(module, 'reset_global_aux_loss_tracker'): module.reset_global_aux_loss_tracker() + if getattr(module, 'qb_beta_accum', None) is not None: + module.qb_beta_accum.zero_() + module.qb_beta_count.zero_() def _update_router_expert_bias( @@ -373,6 +376,48 @@ def _update_router_expert_bias( expert_bias.copy_(updated_expert_bias) +def _update_router_qb_beta( + model: List[torch.nn.Module], + config: TransformerConfig, + dp_cp_group: Optional[torch.distributed.ProcessGroup] = None, +): + """Update the quantile-balancing per-expert bias once per global batch. + + Averages each router's accumulated quantile (qb_beta_accum/qb_beta_count) across + DP, EMA-blends it with the current qb_beta, re-centers, and writes it back. + """ + qb_beta_list = [] + qb_beta_accum_list = [] + qb_beta_count_list = [] + for model_chunk in model: + for module in get_attr_wrapped_model(model_chunk, 'modules')(): + if getattr(module, 'qb_beta_accum', None) is not None and module.training: + qb_beta_list.append(module.qb_beta) + qb_beta_accum_list.append(module.qb_beta_accum) + qb_beta_count_list.append(module.qb_beta_count) + + if len(qb_beta_list) == 0: + return + + stacked_beta = torch.stack(qb_beta_list, dim=0) + local_avg_list = [ + accum / count.clamp(min=1).to(accum.dtype) + for accum, count in zip(qb_beta_accum_list, qb_beta_count_list) + ] + stacked_local_avg = torch.stack(local_avg_list, dim=0) + + torch.distributed.all_reduce( + stacked_local_avg, op=torch.distributed.ReduceOp.AVG, group=dp_cp_group + ) + + ema = config.moe_router_quantile_balancing_ema + stacked_new_beta = ema * stacked_beta + (1.0 - ema) * stacked_local_avg + stacked_new_beta = stacked_new_beta - stacked_new_beta.mean(dim=-1, keepdim=True) + + for qb_beta, new_beta in zip(qb_beta_list, stacked_new_beta): + qb_beta.copy_(new_beta) + + def _allreduce_non_tensor_model_parallel_grads( model: List[torch.nn.Module], config: TransformerConfig, @@ -451,6 +496,75 @@ def _allreduce_non_tensor_model_parallel_grads( _allreduce_layernorm_grads = _allreduce_non_tensor_model_parallel_grads +def _allreduce_replicated_grads_over_gtp_remat_group( + model: List[torch.nn.Module], calculate_per_token_loss: bool = False +): + """Complete the gtp_remat / egtp_remat axis reduction for replicated parameters. + + Replicated (non-gtp-sharded) params have a grad per gtp_remat peer (each from distinct data); + the data-parallel collective only reduced the replicate axis, so the + gtp_remat axis is still missing. How to complete it depends on the loss normalization: + + - ``calculate_per_token_loss=False`` (default): the DP collective produced the 1/replicate mean, + so a MEAN (AVG) over the gtp_remat axis yields the exact full (replicate x gtp) mean, keeping + gradient scaling decoupled from the DP degree. (gtp_remat-sharded params self-average via + their reduce-scatter mean and are skipped here.) + - ``calculate_per_token_loss=True``: DDP applies NO 1/dp scaling; finalize divides every grad by + 1/total_global_tokens (which counts the gtp_remat peers' distinct tokens). The gtp_remat axis + must therefore be SUM-reduced (like the DP axis) — an AVG would shrink each grad by 1/gtp. + + No-op when GTP_remat is inactive (group size <= 1). + """ + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat", "expt_gtp_remat"] + ) + gtp_remat_group = pg_collection.gtp_remat + egtp_remat_group = pg_collection.expt_gtp_remat + + dense_active = gtp_remat_group is not None and gtp_remat_group.size() > 1 + expert_active = egtp_remat_group is not None and egtp_remat_group.size() > 1 + if not dense_active and not expert_active: + return + + dense_params, dense_grads = [], [] + expert_params, expert_grads = [], [] + for model_chunk in model: + for name, param in get_attr_wrapped_model(model_chunk, 'named_parameters')(): + if not param.requires_grad or getattr(param, 'is_gtp_weight_remat', False): + continue # GTP-sharded params: their gtp_remat axis is handled by the RS-mean. + grad_attr = _get_main_grad_attr(param) + grad = getattr(param, grad_attr, None) + if grad is None: + continue + grad = _unshard_if_dtensor(grad) + if getattr(param, 'allreduce', True): + dense_params.append(param) + dense_grads.append(grad.data) + else: + expert_params.append(param) + expert_grads.append(grad.data) + + for params, grads, group in ( + (dense_params, dense_grads, gtp_remat_group), + (expert_params, expert_grads, egtp_remat_group), + ): + if not grads or group is None or group.size() <= 1: + continue + coalesced = _flatten_dense_tensors(grads) + # SUM vs AVG per the loss-normalization regime documented above. + op = ( + torch.distributed.ReduceOp.SUM + if calculate_per_token_loss + else torch.distributed.ReduceOp.AVG + ) + torch.distributed.all_reduce(coalesced, op=op, group=group) + for param, buf, synced in zip(params, grads, _unflatten_dense_tensors(coalesced, grads)): + buf.copy_(synced) + grad_attr = _get_main_grad_attr(param) + orig_grad = getattr(param, grad_attr) + setattr(param, grad_attr, _reshard_if_dtensor(buf, orig_grad)) + + def finalize_model_grads( model: List[torch.nn.Module], num_tokens: Optional[torch.Tensor] = None, @@ -492,7 +606,9 @@ def finalize_model_grads( pp_group = pg_collection.pp embd_group = pg_collection.embd pos_emb_group = pg_collection.pos_embd - dp_cp_group = pg_collection.dp_cp + # Full DP x CP x gtp_remat group: num_tokens (the per-token-loss divisor below) counts the + # gtp_remat peers' distinct tokens. Falls back to replicate dp_cp when gtp is inactive. + dp_cp_group = getattr(pg_collection, 'dp_cp_gtp_remat', None) or pg_collection.dp_cp else: tp_group = parallel_state.get_tensor_model_parallel_group() pp_group = parallel_state.get_pipeline_model_parallel_group() @@ -500,6 +616,14 @@ def finalize_model_grads( pos_emb_group = parallel_state.get_position_embedding_group(check_initialized=False) dp_cp_group = parallel_state.get_data_parallel_group(with_context_parallel=True) + # Fence the current stream against all GTP backward grad work before the DP gradient sync. + if config.gtp_weight_remat_size > 1 or config.expert_gtp_weight_remat_size > 1: + from megatron.core.tensor_parallel.gtp_api import ( + wait_for_gtp_grad_reduction_on_current_stream, + ) + + wait_for_gtp_grad_reduction_on_current_stream() + # All-reduce / reduce-scatter across DP replicas. if config.timers is not None: config.timers('all-grads-sync', log_level=1).start(barrier=config.barrier_with_L1_time) @@ -526,6 +650,9 @@ def finalize_model_grads( barrier=config.barrier_with_L1_time ) _allreduce_non_tensor_model_parallel_grads(model, config, tp_group) + _allreduce_replicated_grads_over_gtp_remat_group( + model, calculate_per_token_loss=config.calculate_per_token_loss + ) if config.timers is not None: config.timers('non-tensor-parallel-grads-all-reduce').stop() @@ -547,6 +674,9 @@ def finalize_model_grads( ) _update_router_expert_bias(model, config, tp_dp_cp_group=tp_dp_cp_group) + if config.moe_router_load_balancing_type == "quantile_balancing": + _update_router_qb_beta(model, config, dp_cp_group=dp_cp_group) + reset_model_temporary_tensors(config, model) # normalize gradients for per-token loss normalization. diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/__init__.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/__init__.py index bc9118598d1..bae27be831c 100644 --- a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/__init__.py +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/__init__.py @@ -16,6 +16,7 @@ from .dbuffer import DBuffer from .fully_shard import fully_shard, microbatch +from .optimizer import fully_shard_optimizer from .placement import Flat, Partial, Placement, Placements, Replicate __all__ = [ @@ -26,5 +27,6 @@ "Placements", "Replicate", "fully_shard", + "fully_shard_optimizer", "microbatch", ] diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/dbuffer.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/dbuffer.py index 9b6e6dc44c3..8381f9a3a5c 100644 --- a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/dbuffer.py +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/dbuffer.py @@ -244,15 +244,24 @@ def distribute_tensors( return buffer def _create_or_validate_out( - self, placements: Iterable[Placement], out: "DBuffer | None" + self, + out: "DBuffer | None", + *, + placements: Iterable[Placement] | None = None, + dtype: torch.dtype | None = None, ) -> "DBuffer": - placements = tuple(placements) + if placements is None: + placements = self.placements + else: + placements = tuple(placements) + if dtype is None: + dtype = self.dtype if out is None: return DBuffer( mesh=self.mesh, placements=placements, tensor_shapes=self.layout.tensor_shapes, - dtype=self.dtype, + dtype=dtype, device=self.device, ) @@ -262,24 +271,18 @@ def _create_or_validate_out( raise ValueError(f"Expected out placements {placements!r}, got {out.placements!r}.") if out.layout != self.layout: raise ValueError(f"Expected out layout {self.layout!r}, got {out.layout!r}.") - if out.dtype != self.dtype: - raise ValueError(f"Expected out dtype {self.dtype}, got {out.dtype}.") + if out.dtype != dtype: + raise ValueError(f"Expected out dtype {dtype}, got {out.dtype}.") if out.device != self.device: raise ValueError(f"Expected out device {self.device}, got {out.device}.") return out - def cast(self, dtype: torch.dtype) -> "DBuffer": + def cast(self, dtype: torch.dtype, *, out: "DBuffer | None" = None) -> "DBuffer": """Return this buffer with the same layout and placements in ``dtype``.""" - if self.dtype == dtype: + if self.dtype == dtype and out is None: return self - destination = DBuffer( - mesh=self.mesh, - placements=self.placements, - tensor_shapes=self.layout.tensor_shapes, - dtype=dtype, - device=self.device, - ) + destination = self._create_or_validate_out(out, dtype=dtype) destination.local_buffer.copy_(self.local_buffer) return destination @@ -289,8 +292,9 @@ def redistribute( """Redistribute this buffer to ``new_placements``. This dispatcher supports the one-axis transitions: - Flat -> Replicate, Partial -> Replicate, Partial -> Flat, and - Replicate -> Flat. Other placement changes are intentionally unsupported. + Flat -> Replicate, Partial -> Replicate, Partial -> Flat, + Replicate -> Flat, and Replicate -> Partial. Other placement changes are + intentionally unsupported. """ new_placements = tuple(new_placements) if len(new_placements) != self.mesh.ndim: @@ -304,7 +308,7 @@ def redistribute( if changed_axis is None: if out is None: return self - out = self._create_or_validate_out(new_placements, out) + out = self._create_or_validate_out(out, placements=new_placements) out.local_buffer.copy_(self.local_buffer) return out @@ -319,6 +323,23 @@ def redistribute( return self.reduce_scatter(axis, new_placement, out=out) if isinstance(old_placement, Replicate) and isinstance(new_placement, Flat): return self.scatter(axis, new_placement, out=out) + if isinstance(old_placement, Replicate) and isinstance(new_placement, Partial): + # Replicate and Partial share the same local layout, so relabel the + # buffer without communication. Value-preserving for AVG only -- the + # mean of identical per-rank locals is that value; SUM would need a + # 1/axis_size rescale, which no caller needs. + if new_placement.reduce_op != dist.ReduceOp.AVG: + raise NotImplementedError( + "Replicate -> Partial redistribute supports AVG only, got " + f"{new_placement.reduce_op!r}." + ) + if out is not None: + raise NotImplementedError( + "Replicate -> Partial redistribute does not support an out buffer." + ) + return DBuffer.from_local( + self.local_buffer, self.mesh, new_placements, self.layout.tensor_shapes + ) raise NotImplementedError( "Unsupported DBuffer placement transition on axis " f"{axis}: {old_placement!r} -> {new_placement!r}." @@ -334,7 +355,7 @@ def allgather(self, mesh_axis: int, *, out: "DBuffer | None" = None) -> "DBuffer placements = list(self.placements) placements[mesh_axis] = Replicate() _validate_placements(placements) - out = self._create_or_validate_out(placements, out) + out = self._create_or_validate_out(out, placements=placements) dist.all_gather_into_tensor( output_tensor=out.local_buffer, input_tensor=self.local_buffer, @@ -351,7 +372,7 @@ def allreduce(self, mesh_axis: int, *, out: "DBuffer | None" = None) -> "DBuffer placements = list(self.placements) placements[axis] = Replicate() - out = self._create_or_validate_out(placements, out) + out = self._create_or_validate_out(out, placements=placements) out.local_buffer.copy_(self.local_buffer) dist.all_reduce( out.local_buffer, op=partial_placement.reduce_op, group=self.mesh.get_group(axis) @@ -372,7 +393,7 @@ def reduce_scatter( placements = list(self.placements) placements[axis] = new_placement _validate_placements(placements) - out = self._create_or_validate_out(placements, out) + out = self._create_or_validate_out(out, placements=placements) dist.reduce_scatter_tensor( output=out.local_buffer, input=self.local_buffer, @@ -400,7 +421,7 @@ def scatter( self.mesh, placements ) else: - out = self._create_or_validate_out(placements, out) + out = self._create_or_validate_out(out, placements=placements) destination_offset = out.offset destination_numel = out.local_buffer.numel() @@ -449,7 +470,17 @@ def get_dtensor(self, index: int) -> DTensor: elif isinstance(placement, Flat): torch_placements.append(dist_tensor.Shard(0)) elif isinstance(placement, Partial): - raise ValueError("Partial DBuffer placements cannot be represented as DTensor.") + # main_grad backs .grad while it rests DP-outer-Partial between + # microbatches, so a Partial placement must round-trip to a DTensor. + if placement.reduce_op == dist.ReduceOp.AVG: + reduce_op = "avg" + elif placement.reduce_op == dist.ReduceOp.SUM: + reduce_op = "sum" + else: + raise ValueError( + f"Unsupported Partial reduce op for DTensor: {placement.reduce_op!r}." + ) + torch_placements.append(dist_tensor.Partial(reduce_op)) else: raise TypeError(f"Unsupported placement for DTensor conversion: {placement!r}.") diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/fully_shard.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/fully_shard.py index 0ab256a7ef4..0f53c34f359 100644 --- a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/fully_shard.py +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/fully_shard.py @@ -32,13 +32,13 @@ def fully_shard( mixed_precision_policy: MixedPrecisionPolicy | None = None, use_symm_mem: bool = False, ) -> None: - """Shard one module as a per-module FSDP unit. + """Apply FSDP to a module in place. This attaches the FSDP mixin to the original module instance, so parent modules do not need to replace existing child-module references. Args: - module: Module whose currently unowned parameters become this FSDP unit. + module: Module whose currently unowned parameters are managed by FSDP. mesh: Device mesh used for sharding. placements: Parameter, gradient, and optimizer placements. mixed_precision_policy: Optional precision policy. Defaults to FP32 main weights @@ -68,10 +68,14 @@ def fully_shard( @contextmanager def microbatch(module: nn.Module, is_last: bool) -> Iterator[None]: - """Scope experimental FSDP state to one microbatch. + """Mark an FSDP microbatch as the last accumulation microbatch. + + At present, this is only needed for HSDP/HFSDP gradient accumulation, so + FSDP finalizes gradients only on the last backward. Plain all-Flat data + parallelism finalizes gradients on every backward and does not need it. Args: - module: Module tree whose experimental FSDP roots should use this microbatch state. + module: Module tree whose FSDP roots should use this microbatch state. is_last: Whether forwards in this scope are for the last microbatch. """ contexts: list[FsdpContext] = [] diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/indexed_order.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/indexed_order.py new file mode 100644 index 00000000000..d7c9c63ed0c --- /dev/null +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/indexed_order.py @@ -0,0 +1,53 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Ordered sequence with indexed item lookup.""" + +from collections.abc import Iterator +from typing import Generic, TypeVar + +T = TypeVar("T") + + +class IndexedOrder(Generic[T]): + """Insertion order with constant-time successor lookup by item.""" + + def __init__(self) -> None: + """Create an empty indexed order.""" + self._items: list[T] = [] + self._index_by_item: dict[T, int] = {} + + def append(self, item: T) -> None: + """Append ``item`` to the order. + + Args: + item: Item to append. + + Raises: + ValueError: If ``item`` is already present in the order. + """ + if item in self._index_by_item: + raise ValueError("IndexedOrder does not support duplicate items.") + self._index_by_item[item] = len(self._items) + self._items.append(item) + + def __iter__(self) -> Iterator[T]: + """Iterate over items in order.""" + return iter(self._items) + + def next_item(self, item: T) -> T | None: + """Return the item that follows ``item``, if any.""" + index = self._index_by_item[item] + next_index = index + 1 + return self._items[next_index] if next_index < len(self._items) else None diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/layout.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/layout.py index 11070495e4d..e775e595cc4 100644 --- a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/layout.py +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/layout.py @@ -88,7 +88,7 @@ def build(cls, shapes: Iterable[Shape], dp_size: int) -> "GlobalLayout": ) chunk_size = math.lcm(chunk_size, row_size) - # chunk_size is the packing unit. Since every tensor row size divides it, + # chunk_size is the packing granularity. Since every tensor row size divides it, # DP shard boundaries that are multiples of chunk_size avoid splitting dim-0 rows. UNASSIGNED_OFFSET = -1 tensor_to_offset: list[int] = [UNASSIGNED_OFFSET] * len(tensor_shapes) diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/module.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/module.py index 3c97fda3242..33e6985d8ed 100644 --- a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/module.py +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/module.py @@ -14,8 +14,6 @@ """Module mixin for the minimal Megatron-FSDP path.""" -import dataclasses -from collections import deque from collections.abc import Callable from typing import Literal, cast @@ -24,28 +22,26 @@ from torch.distributed import DeviceMesh from ..mixed_precision import MixedPrecisionPolicy -from .parameter_group import FsdpParameterGroup, contained_in_parameter_group +from .indexed_order import IndexedOrder +from .parameter_group import FsdpParameterGroup, get_containing_parameter_group from .placement import MeshAxis, Placements -@dataclasses.dataclass(frozen=True) -class DelayedRelease: - """A module whose unsharded storage can be released after its consumer event.""" - - consumer_event: torch.cuda.Event | None - module: "FsdpModule" - - class FsdpContext: - """Runtime state, stream, and release scheduler shared by one FSDP subtree.""" + """Runtime stream and prefetch state shared by one FSDP subtree.""" allgather_stream: torch.cuda.Stream - delayed_releases: deque[DelayedRelease] + reduce_scatter_stream: torch.cuda.Stream # HFSDP/HSDP need explicit last-microbatch state. First-microbatch state is # unnecessary because it can be detected when ``model_weight``, after syncing # from ``main_weight``, has placements different from ``Placements.optimizer``. is_last_microbatch: bool root_module: "FsdpModule" + # Static orders used to drive all-gather prefetch. We may want to switch to + # capturing runtime order if static module order proves too fragile. Each + # FsdpModule tracks its own materialized state via ``FsdpModule._unshard_event``. + forward_order: IndexedOrder["FsdpModule"] + backward_order: IndexedOrder["FsdpModule"] def __init__(self, device: torch.device, root_module: "FsdpModule") -> None: """Create rank-local runtime state for a root FSDP subtree. @@ -56,26 +52,29 @@ def __init__(self, device: torch.device, root_module: "FsdpModule") -> None: """ self.root_module = root_module self.is_last_microbatch = True - self.delayed_releases = deque() + self.forward_order = IndexedOrder() + self.backward_order = IndexedOrder() with torch.cuda.device(device): self.allgather_stream = torch.cuda.Stream() + self.reduce_scatter_stream = torch.cuda.Stream() + + def current_stream(self) -> torch.cuda.Stream: + """Current stream on this context's device.""" + return torch.cuda.current_stream(self.allgather_stream.device) + + def register_post_backward_final_callback(self) -> None: + """Register this root context's final callback for the current backward. - def enqueue_release(self, module: "FsdpModule") -> None: - """Queue a module's unsharded storage for delayed release.""" - consumer_event = torch.cuda.current_stream(self.allgather_stream.device).record_event() - self.delayed_releases.append(DelayedRelease(consumer_event=consumer_event, module=module)) + Root ``post_backward()`` means only that root-owned parameters have + accumulated gradients; it may run before descendant reductions, or not + run at all when the root owns no trainable parameters. Waiting at + autograd completion orders consumers after every descendant reduction. + """ - def drain_delayed_releases(self, target_length: int) -> None: - """Release queued module storages FIFO until the queue reaches ``target_length``.""" - if target_length < 0: - raise ValueError(f"target_length must be non-negative, got {target_length}.") + def post_backward_final_callback() -> None: + self.current_stream().wait_stream(self.reduce_scatter_stream) - while len(self.delayed_releases) > target_length: - delayed_release = self.delayed_releases.popleft() - with torch.cuda.stream(self.allgather_stream): - if delayed_release.consumer_event is not None: - self.allgather_stream.wait_event(delayed_release.consumer_event) - delayed_release.module.release_unsharded_storage() + torch.autograd.Variable._execution_engine.queue_callback(post_backward_final_callback) class FsdpModule: @@ -87,7 +86,11 @@ class FsdpModule: _parameter_groups: tuple[FsdpParameterGroup, ...] _context: FsdpContext | None _ready_grad_parameters: set[nn.Parameter] - _num_training_parameters: int + _num_trainable_parameters: int + # Event recorded after this FsdpModule's full parameters are materialized. + # ``None`` lets pre_forward enqueue an all-gather unless an earlier FsdpModule + # already prefetched this module. + _unshard_event: torch.cuda.Event | None def __init__( self, @@ -99,6 +102,7 @@ def __init__( """Initialize FSDP runtime state on an already-constructed module.""" self._context = None self._name = None + self._unshard_event = None owned_parameters = _collect_owned_parameters(self) axis_indices = tuple(_axis_index(mesh, axis) for axis in placements.dp_axes) assert axis_indices == tuple( @@ -117,7 +121,7 @@ def __init__( ] self._parameter_groups = tuple(parameter_groups) self._ready_grad_parameters = set() - self._num_training_parameters = sum( + self._num_trainable_parameters = sum( len(group.sharded_parameters) for group in self._parameter_groups if group.requires_grad ) self._register_hooks() @@ -147,8 +151,15 @@ def _lazy_init_context(self) -> None: if self._context is not None: return - context = FsdpContext(device=self._parameter_groups[0].main_weight.device, root_module=self) - for submodule_name, submodule in cast(nn.Module, self).named_modules(): + root_module = cast(nn.Module, self) + first_parameter = next(root_module.parameters(), None) + if first_parameter is None: + raise RuntimeError("FSDP root module requires at least one parameter in its subtree.") + + context = FsdpContext(device=first_parameter.device, root_module=self) + # named_modules() yields FsdpModules in registration order, which is the static + # forward execution order used to prefetch the next FsdpModule's all-gather. + for submodule_name, submodule in root_module.named_modules(): if not isinstance(submodule, FsdpModule): continue if submodule._context is not None: @@ -158,6 +169,11 @@ def _lazy_init_context(self) -> None: ) submodule._context = context submodule._name = submodule_name + context.forward_order.append(submodule) + + # Backward starts from the root pre-backward hook before visiting child + # subtrees in reverse module order. + _collect_backward_order(root_module, context.backward_order) @property def context(self) -> FsdpContext: @@ -167,14 +183,14 @@ def context(self) -> FsdpContext: @property def name(self) -> str: - """Return this FSDP unit's name.""" + """Return this FsdpModule's name.""" name = self._name if name is None: raise RuntimeError("FSDP module name has not been initialized.") return name def is_root(self) -> bool: - """Return whether this module is the outermost FSDP unit in its context.""" + """Return whether this module is the outermost FsdpModule in its context.""" return self.context.root_module is self def _register_hooks(self) -> None: @@ -182,10 +198,16 @@ def _register_hooks(self) -> None: module.register_forward_pre_hook(lambda _module, _args: self.pre_forward()) module.register_forward_hook(lambda _module, _args, _output: self.post_forward()) module.register_full_backward_pre_hook(lambda _module, _grad_output: self.pre_backward()) - # Gradient reduction is parameter-completion based: once every owned - # Parameter has accumulated its grad, this FSDP unit can reduce and - # reshard. Module full-backward hooks can fire before that when module - # inputs do not require grad. + if self._num_trainable_parameters == 0: + module.register_full_backward_hook( + lambda _module, _grad_input, _grad_output: self.post_backward() + ) + return + + # Gradient reduction for trainable parameters is parameter-completion + # based: once every owned Parameter has accumulated its grad, this + # FsdpModule can reduce and reshard. Module full-backward hooks can fire + # before that when module inputs do not require grad. for group in self._parameter_groups: if not group.requires_grad: continue @@ -195,74 +217,131 @@ def _register_hooks(self) -> None: def _make_grad_hook(self, parameter: nn.Parameter) -> Callable[[nn.Parameter], None]: def grad_hook(_parameter: nn.Parameter) -> None: self._ready_grad_parameters.add(parameter) - if len(self._ready_grad_parameters) == self._num_training_parameters: + if len(self._ready_grad_parameters) == self._num_trainable_parameters: self.post_backward() return grad_hook def pre_forward(self) -> None: - """Prepare full parameters for forward compute.""" + """Prepare full parameters for forward compute and prefetch the next FsdpModule. + + While this FsdpModule computes, we issue the next FsdpModule's all-gather + on the comm stream, so ``AG_{i+1}`` is launched before ``F_i`` finishes. + """ self._lazy_init_context() torch.cuda.nvtx.range_push(self._nvtx_label("forward")) self._ready_grad_parameters.clear() + context = self.context + allgather_stream = context.allgather_stream + current_stream = context.current_stream() + if self.is_root(): - allgather_stream = self.context.allgather_stream - allgather_stream.wait_stream(torch.cuda.current_stream(allgather_stream.device)) - self._unshard_parameter_groups(sync_model_weight=True) + allgather_stream.wait_stream(current_stream) - def _unshard_parameter_groups(self, *, sync_model_weight: bool) -> None: - """Materialize full parameters for this FSDP unit.""" - self.context.drain_delayed_releases(target_length=1) + self._unshard_parameter_groups() + assert self._unshard_event is not None + # Compute waits only for this FsdpModule's all-gather (the prefetch below is + # issued afterwards, so it is free to run concurrently with this FsdpModule). + current_stream.wait_event(self._unshard_event) - allgather_stream = self.context.allgather_stream - current_stream = torch.cuda.current_stream(allgather_stream.device) + next_module = context.forward_order.next_item(self) + if next_module is not None: + next_module._unshard_parameter_groups() + def _unshard_parameter_groups(self) -> None: + """Unshard this FsdpModule's parameter groups on the all-gather stream. + + If ``_unshard_event`` is already set, this FsdpModule was already + unsharded or prefetched and this method is a no-op. Otherwise, this + method records ``_unshard_event`` after materialization so compute + can wait without depending on later release work. + """ + if self._unshard_event is not None: + return + + allgather_stream = self.context.allgather_stream with torch.cuda.stream(allgather_stream): for group in self._parameter_groups: - if sync_model_weight: - # TODO: After NVIDIA/Megatron-LM#5411 lands, move this sync to the - # optimizer post-step hook instead of running it every microbatch. - group.sync_model_weight_from_main_weight() group.unshard_parameters() - current_stream.wait_stream(allgather_stream) + self._unshard_event = allgather_stream.record_event() def post_forward(self) -> None: """Return parameters to their sharded resting state after forward compute.""" self._reshard_parameter_groups() - self.context.enqueue_release(self) - if self.is_root(): - self.context.drain_delayed_releases(target_length=0) torch.cuda.nvtx.range_pop() def _reshard_parameter_groups(self) -> None: + """Reshard parameter groups and release unsharded storage after compute. + + This method clears ``_unshard_event`` after queuing the release, so + future users enqueue a fresh all-gather. + """ for group in self._parameter_groups: group.reshard_parameters() + allgather_stream = self.context.allgather_stream + allgather_stream.wait_stream(self.context.current_stream()) + # Release on the all-gather stream where unsharded storage was allocated, + # so no record_stream() call is required for the storage. + with torch.cuda.stream(allgather_stream): + for group in self._parameter_groups: + group.release_unsharded_storage() + self._unshard_event = None + def pre_backward(self) -> None: - """Prepare full parameters for backward compute.""" + """Prepare full parameters and prefetch the next FsdpModule in backward order.""" torch.cuda.nvtx.range_push(self._nvtx_label("backward")) - self._unshard_parameter_groups(sync_model_weight=False) + context = self.context + current_stream = context.current_stream() + if self.is_root(): + context.register_post_backward_final_callback() + # Fork the reduce-scatter stream from the current stream once, at the + # start of backward, so every module's post-backward reduce-scatter is + # part of any active CUDA-graph capture. A stream only joins the + # capture via this wait_stream edge; without it the first allocation on + # the reduce-scatter stream falls back to a raw cudaMalloc, which is + # illegal during capture. Later modules are covered by the post-copy + # fork each preceding module issues before its collective. + context.reduce_scatter_stream.wait_stream(current_stream) + + self._unshard_parameter_groups() + assert self._unshard_event is not None + current_stream.wait_event(self._unshard_event) + + next_module = context.backward_order.next_item(self) + if next_module is not None: + next_module._unshard_parameter_groups() def post_backward(self) -> None: """Reduce gradients and return parameters to their sharded resting state.""" - for group in self._parameter_groups: - if group.requires_grad: - group.reduce_gradients() + self._reduce_gradient_groups() self._reshard_parameter_groups() - self.context.enqueue_release(self) - if self.is_root(): - self.context.drain_delayed_releases(target_length=0) self._ready_grad_parameters.clear() torch.cuda.nvtx.range_pop() - def release_unsharded_storage(self) -> None: - """Release unsharded storage owned by this FSDP unit.""" + def _reduce_gradient_groups(self) -> None: + """Pack gradients and immediately launch their reduce-scatters.""" + context = self.context + reduce_scatter_stream = context.reduce_scatter_stream + current_stream = context.current_stream() + for group in self._parameter_groups: - group.release_unsharded_storage() + if not group.requires_grad: + continue + + with torch.cuda.stream(reduce_scatter_stream): + partial_grad = group.allocate_partial_grad_buffer() + + current_stream.wait_stream(reduce_scatter_stream) + group.copy_gradients_to_partial_buffer(partial_grad) + + reduce_scatter_stream.wait_stream(current_stream) + with torch.cuda.stream(reduce_scatter_stream): + group.reduce_partial_gradients(partial_grad, self.context.is_last_microbatch) @property def parameter_groups(self) -> tuple[FsdpParameterGroup, ...]: - """Parameter groups owned by this FSDP unit.""" + """Parameter groups owned by this FsdpModule.""" return self._parameter_groups def _nvtx_label(self, phase: Literal["forward", "backward"]) -> str: @@ -270,6 +349,15 @@ def _nvtx_label(self, phase: Literal["forward", "backward"]) -> str: return f"MFSDP {name} {phase}" +def _collect_backward_order(module: nn.Module, order: IndexedOrder["FsdpModule"]) -> None: + """Collect FsdpModules in static backward prefetch order.""" + if isinstance(module, FsdpModule): + order.append(module) + + for child in reversed(list(module.children())): + _collect_backward_order(child, order) + + def _axis_index(mesh: DeviceMesh, axis: MeshAxis) -> int: if isinstance(axis, int): axis_index = axis @@ -295,8 +383,10 @@ def visit(submodule: nn.Module, submodule_fqn: str) -> None: parameter_fqn = ( f"{submodule_fqn}.{local_parameter_name}" if submodule_fqn else local_parameter_name ) - if contained_in_parameter_group(parameter): - raise ValueError(f"Parameter {parameter_fqn!r} is already owned by an FSDP unit.") + if get_containing_parameter_group(parameter) is not None: + raise ValueError( + f"Parameter {parameter_fqn!r} is already owned by another FsdpModule." + ) parameters[parameter_fqn] = parameter for child_name, child_module in submodule.named_children(): @@ -306,8 +396,6 @@ def visit(submodule: nn.Module, submodule_fqn: str) -> None: visit(child_module, child_fqn) visit(root_module, "") - if not parameters: - raise ValueError("fully_shard requires at least one unowned parameter.") return parameters diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/optimizer.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/optimizer.py new file mode 100644 index 00000000000..f1617141569 --- /dev/null +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/optimizer.py @@ -0,0 +1,124 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Optimizer adapter for the minimal Megatron-FSDP path.""" + +from typing import Any, NamedTuple + +import torch +from torch import nn + +from .parameter_group import FsdpParameterGroup, get_containing_parameter_group + + +def fully_shard_optimizer( + optimizer: torch.optim.Optimizer, *, precision_aware: bool = False +) -> None: + """Attach FSDP-aware step hooks to an optimizer instance. + + The adapted optimizer preserves its existing parameter groups, temporarily + casts gradients around optimizer steps for FSDP sharded parameters whose + data dtype differs from their grad dtype unless the optimizer is precision + aware, and refreshes compute weights after each optimizer step. + + Alternatives considered: + - Monkey-patching optimizer methods directly on the instance. This is + more invasive and harder to compose than hooks. + - Generating an FSDP-specific subclass per ``torch.optim.Optimizer``. + This adds extra class-generation machinery, but would let us + instrument ``zero_grad`` and ``__init__`` as well as ``step`` if needed. + - Casting from ``main_grad.dtype`` to ``main_weight.dtype`` after the + last microbatch and casting back before the first microbatch. This + should be done from a root post-backward callback if needed later, so + users do not need to call ``fully_shard_optimizer`` on an existing + ``torch.optim.Optimizer``. + - Letting the user set ``main_weight`` and ``main_grad`` to the same + dtype. This is enough for an FSDP2 drop-in replacement path and lets + optimizers stay unaware of FSDP precision handling. + + Args: + optimizer: Optimizer instance to adapt in place. + precision_aware: Whether the optimizer accepts FSDP's mixed-precision + gradients without temporary casting. + """ + + class CastedGrad(NamedTuple): + """Original grad tensor temporarily replaced during an optimizer step.""" + + parameter: nn.Parameter + original_grad: torch.Tensor + + def set_grad(parameter: nn.Parameter, grad: torch.Tensor) -> None: + """Install a grad with matching grad_dtype on a sharded parameter.""" + # Clear the existing grad before switching grad_dtype; the sharded + # parameter cannot advertise a new grad dtype while the old grad + # object with the previous dtype is still attached. + parameter.grad = None + parameter.grad_dtype = grad.dtype + parameter.grad = grad + + casted_grads: list[CastedGrad] = [] + + def step_pre_hook( + hooked_optimizer: torch.optim.Optimizer, args: tuple[Any, ...], kwargs: dict[str, Any] + ) -> None: + closure = kwargs.get("closure") + if closure is None and len(args) > 1: + closure = args[1] + if closure is not None: + # Step hooks run outside the base optimizer step, but closures run inside it. + # We need to cast grads after the closure materializes them and before the + # optimizer consumes them, which this hook-only adapter cannot intercept. + raise NotImplementedError( + "fully_shard_optimizer does not support optimizer.step closures." + ) + assert not casted_grads + for group in hooked_optimizer.param_groups: + for parameter in group["params"]: + if not isinstance(parameter, nn.Parameter): + raise TypeError( + "fully_shard_optimizer expected optimizer param groups to contain " + f"nn.Parameter values, got {type(parameter)!r}." + ) + if precision_aware or get_containing_parameter_group(parameter) is None: + continue + if parameter.grad is None: + continue + if parameter.grad.dtype == parameter.dtype: + continue + + casted_grads.append(CastedGrad(parameter, parameter.grad)) + set_grad(parameter, parameter.grad.to(dtype=parameter.dtype)) + + def step_post_hook( + hooked_optimizer: torch.optim.Optimizer, args: tuple[Any, ...], kwargs: dict[str, Any] + ) -> None: + del args, kwargs + for parameter, original_grad in casted_grads: + set_grad(parameter, original_grad) + casted_grads.clear() + + fsdp_parameter_groups: set[FsdpParameterGroup] = set() + for optimizer_group in hooked_optimizer.param_groups: + for parameter in optimizer_group["params"]: + parameter_group = get_containing_parameter_group(parameter) + if parameter_group is None: + continue + fsdp_parameter_groups.add(parameter_group) + + for parameter_group in fsdp_parameter_groups: + parameter_group.sync_model_weight_from_main_weight() + + optimizer.register_step_pre_hook(step_pre_hook) + optimizer.register_step_post_hook(step_post_hook) diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/parameter_group.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/parameter_group.py index eeec848416b..fd5eb5d2033 100644 --- a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/parameter_group.py +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/parameter_group.py @@ -14,7 +14,6 @@ """Parameter-group runtime state for the minimal Megatron-FSDP path.""" -from collections.abc import Iterable from contextlib import nullcontext import torch @@ -30,9 +29,9 @@ _CONTAINING_PARAMETER_GROUP_ATTR = "_mfsdp_parameter_group" -def contained_in_parameter_group(parameter: nn.Parameter) -> bool: - """Return whether a parameter is already owned by an FsdpParameterGroup.""" - return hasattr(parameter, _CONTAINING_PARAMETER_GROUP_ATTR) +def get_containing_parameter_group(parameter: nn.Parameter) -> "FsdpParameterGroup | None": + """Return the FSDP parameter group that owns ``parameter``, if any.""" + return getattr(parameter, _CONTAINING_PARAMETER_GROUP_ATTR, None) class FsdpParameterGroup: @@ -151,13 +150,9 @@ def __init__( "main_grad is built from main_weight tensor shapes on the same mesh, " "and DBuffer layouts are deterministic from those shapes and mesh size." ) - if self.main_grad.placements != self.main_weight.placements: - raise ValueError( - "FSDP temporarily requires main_grad and main_weight to have the same " - "placements until HSDP/HFSDP support is implemented. " - f"Got main_grad placements {self.main_grad.placements} and " - f"main_weight placements {self.main_weight.placements}." - ) + # main_grad rests here (DP-outer-Partial for HSDP) between microbatches and + # is finalized to main_weight's placements after the last microbatch. + self._accumulation_placements = main_grad_placements sharded_parameters: list[nn.Parameter] = [] unsharded_parameters: list[nn.Parameter] = [] main_grad_dtype = self.main_grad.dtype if self.main_grad is not None else None @@ -177,6 +172,9 @@ def __init__( self.sharded_parameters = tuple(sharded_parameters) self.unsharded_parameters = tuple(unsharded_parameters) + # Compute weights must be initialized before the first forward; subsequent + # refreshes happen from the FSDP optimizer's post-step hook. + self.sync_model_weight_from_main_weight() self._switch_to_sharded_parameters() self._unsharded_model_weight.release_storage() @@ -201,6 +199,13 @@ def sync_model_weight_from_main_weight(self) -> None: if self.main_weight is self.model_weight: return + if self.main_weight.placements == self.model_weight.placements: + self.main_weight.cast(self.model_weight.dtype, out=self.model_weight) + return + + # main_weight is typically the higher-precision optimizer dtype, while + # model_weight is the lower-precision compute dtype. Cast before redistributing + # so cross-rank communication moves the smaller compute-dtype payload. self.main_weight.cast(self.model_weight.dtype).redistribute( self.model_weight.placements, out=self.model_weight ) @@ -245,11 +250,55 @@ def release_unsharded_storage(self) -> None: # so keep the shared storage-release path. self._unsharded_model_weight.release_storage() - def reduce_gradients(self) -> None: - """Reduce full local gradients into sharded parameter gradients.""" + def _install_sharded_grads(self) -> None: + """Point each sharded parameter's grad at main_grad's current DTensor view.""" + assert self.main_grad is not None + for index, sharded_parameter in enumerate(self.sharded_parameters): + sharded_parameter.grad = self.main_grad.get_dtensor(index) + + def allocate_partial_grad_buffer(self) -> DBuffer: + """Allocate the unreduced reduce-scatter input buffer.""" assert self.main_grad is not None - def has_grad(parameters: Iterable[nn.Parameter]) -> bool: + # NCCL symmetric-memory reduce-scatter only selects the symmetric kernel for SUM today. + # Preserve AVG semantics by reducing SUM and scaling the output below. + partial_op = dist.ReduceOp.AVG if self._symm_mem_pool is None else dist.ReduceOp.SUM + grads: list[torch.Tensor] = [] + for name, parameter in zip(self.parameter_names, self.unsharded_parameters, strict=True): + if parameter.grad is None: + raise RuntimeError(f"Missing gradient for FSDP parameter {name!r}.") + grads.append(parameter.grad) + with self._symmetric_memory_context(): + return DBuffer( + mesh=self.mesh, + placements=[Partial(partial_op)] * self.mesh.ndim, + tensor_shapes=tuple(grad.shape for grad in grads), + dtype=grads[0].dtype, + device=grads[0].device, + ) + + def copy_gradients_to_partial_buffer(self, partial_grad: DBuffer) -> None: + """Pack full local gradients into an existing reduce-scatter input buffer.""" + # A future fused-wgrad path can write directly into these buffer views. + for index, parameter in enumerate(self.unsharded_parameters): + partial_grad.get_local_tensor(index).copy_(parameter.grad) + parameter.grad = None + + def reduce_partial_gradients( + self, partial_grad: DBuffer, is_last_microbatch: bool = True + ) -> None: + """Reduce a packed partial gradient buffer into sharded parameter gradients. + + For HSDP main_grad rests DP-outer-Partial (Partial where main_weight is + Replicate) between microbatches, accumulating each backward through the + standard zero_grad contract; the last microbatch reduces the DP-outer axes, + finalizing main_grad to main_weight's placements so ``.grad`` is the fully + reduced gradient before ``optimizer.step()``. With every axis Flat (plain + DP) main_grad already rests finalized. + """ + assert self.main_grad is not None + + def has_grad(parameters: tuple[nn.Parameter, ...]) -> bool: has_any_grad = False has_any_missing_grad = False for parameter in parameters: @@ -261,31 +310,28 @@ def has_grad(parameters: Iterable[nn.Parameter]) -> bool: raise RuntimeError("FSDP sharded gradients must be either all set or all None.") return has_any_grad - grads: list[torch.Tensor] = [] - for name, parameter in zip(self.parameter_names, self.unsharded_parameters, strict=True): - if parameter.grad is None: - raise RuntimeError(f"Missing gradient for FSDP parameter {name!r}.") - grads.append(parameter.grad) - - # NCCL symmetric-memory reduce-scatter only selects the symmetric kernel for SUM today. - # Preserve AVG semantics by reducing SUM and scaling the output below. - partial_op = dist.ReduceOp.AVG if self._symm_mem_pool is None else dist.ReduceOp.SUM - with self._symmetric_memory_context(): - partial_grad = DBuffer.distribute_tensors( - grads, mesh=self.mesh, placements=[Partial(partial_op)] * self.mesh.ndim - ) - - # zero_grad(set_to_none=True) clears sharded parameter grads, so the next + # zero_grad(set_to_none=True) clears sharded parameter grads, so this # backward can reduce directly into main_grad. zero_grad(set_to_none=False) # leaves sharded grads installed, so this backward accumulates into main_grad. has_sharded_grads = has_grad(self.sharded_parameters) + + # A non-accumulation main_grad means the previous step finalized it; this + # only happens on the first microbatch. Redistribute it back to the + # DP-outer-Partial accumulation placement -- a metadata relabel for HSDP, + # and a fresh reduce-scattered buffer for HFSDP in the future. + if self.main_grad.placements != self._accumulation_placements: + self.main_grad = self.main_grad.redistribute(self._accumulation_placements) + if has_sharded_grads: + self._install_sharded_grads() + can_reduce_into_main_grad = ( not has_sharded_grads and partial_grad.dtype == self.main_grad.dtype ) reduce_axis = changed_mesh_axis(partial_grad.placements, self.main_grad.placements) if reduce_axis is None: raise RuntimeError("FSDP gradient reduction requires a changed placement axis.") - grad_divisor = self.mesh.size(reduce_axis) if partial_op == dist.ReduceOp.SUM else 1 + partial_reduce_op = partial_grad.placements[reduce_axis].reduce_op + grad_divisor = self.mesh.size(reduce_axis) if partial_reduce_op == dist.ReduceOp.SUM else 1 if self._symm_mem_pool is not None: partial_grad.rendezvous(reduce_axis) if can_reduce_into_main_grad: @@ -300,13 +346,14 @@ def has_grad(parameters: Iterable[nn.Parameter]) -> bool: self.main_grad.local_buffer.add_(reduced_grad.local_buffer) else: self.main_grad.local_buffer.copy_(reduced_grad.local_buffer) - if not has_sharded_grads: - for index, parameter in enumerate(self.sharded_parameters): - parameter.grad = self.main_grad.get_dtensor(index) + self._install_sharded_grads() - for parameter in self.unsharded_parameters: - parameter.grad = None + if is_last_microbatch: + # Finalize the deferred DP-outer reduction (all-reduce for HSDP, + # reduce-scatter for HFSDP) and install the sharded parameter gradients. + self.main_grad = self.main_grad.redistribute(self.main_weight.placements) + self._install_sharded_grads() def _get_parameter_owner(module: nn.Module, name: str) -> tuple[nn.Module, str]: diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/placement.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/placement.py index 75d3af4368c..6f99ff1a06c 100644 --- a/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/placement.py +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/experimental/placement.py @@ -56,7 +56,7 @@ class Partial(Placement): @dataclasses.dataclass(frozen=True) class Flat(Placement): - """Flat per-unit dim-0 sharded local buffer placement.""" + """Flat dim-0 sharded local buffer placement.""" def changed_mesh_axis( diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/megatron_fsdp.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/megatron_fsdp.py index 11b0768ad9d..bb6ddc24500 100644 --- a/megatron/core/distributed/fsdp/src/megatron_fsdp/megatron_fsdp.py +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/megatron_fsdp.py @@ -1498,7 +1498,7 @@ def forward(self, *inputs, **kwargs): self._replace_param_with_raw_if_needed() with torch.autograd.profiler.record_function("CustomFSDP.forward"): # Call the forward pass of the wrapped module. - output = self.module.forward(*inputs, **kwargs) + output = self.module(*inputs, **kwargs) return output diff --git a/megatron/core/distributed/fsdp/src/megatron_fsdp/package_info.py b/megatron/core/distributed/fsdp/src/megatron_fsdp/package_info.py index c6f2355eddb..d5f083f665f 100644 --- a/megatron/core/distributed/fsdp/src/megatron_fsdp/package_info.py +++ b/megatron/core/distributed/fsdp/src/megatron_fsdp/package_info.py @@ -2,7 +2,7 @@ MAJOR = 0 -MINOR = 6 +MINOR = 7 PATCH = 0 PRE_RELEASE = 'rc0' diff --git a/megatron/core/distributed/param_and_grad_buffer.py b/megatron/core/distributed/param_and_grad_buffer.py index 691affe6be6..8e490f8c539 100644 --- a/megatron/core/distributed/param_and_grad_buffer.py +++ b/megatron/core/distributed/param_and_grad_buffer.py @@ -31,7 +31,7 @@ from ..fp8_utils import ( _stage_param_to_bf16, copy_back_gathered_bf16_into_fp8_param, - copy_tensor_to_quantized_param, + copy_tensors_to_quantized_params, is_float8tensor, is_grouped_mxfp8tensor, is_grouped_tensor, @@ -379,6 +379,9 @@ def _post_param_sync(self): if bucket.param_data is None: continue has_non_quantized_weight = False + quantized_params = [] + param_slices = [] + flat_param_data = bucket.param_data.view(-1) for param in bucket.params: # Non-quantized weights are already mapped to param.data. Skip # mixed buckets because zeroing bucket.param_data would also @@ -387,8 +390,11 @@ def _post_param_sync(self): has_non_quantized_weight = True break param_start, param_end = bucket.param_to_index[param] - param_slice = bucket.param_data.view(-1)[param_start:param_end] - copy_tensor_to_quantized_param(param, param_slice) + quantized_params.append(param) + param_slices.append(flat_param_data[param_start:param_end]) + # Cast the bucket in one call: these casts are small, so the per-param cost of + # issuing them is worth avoiding. + copy_tensors_to_quantized_params(quantized_params, param_slices) if has_non_quantized_weight: continue # All-gathered params are not needed after being copied to param.data. @@ -812,7 +818,12 @@ def start_grad_sync(self, force_all_reduce: Optional[bool] = False): ) if async_op: - if self.ddp_config.reduce_scatter_with_fp32_accumulation and not force_all_reduce: + # fp32-accum RS needs the distributed optimizer; else fall through (all-reduce -> cm). + if ( + self.ddp_config.reduce_scatter_with_fp32_accumulation + and self.ddp_config.use_distributed_optimizer + and not force_all_reduce + ): assert ( len(self.buckets) == 1 ), "Only 1 bucket supported with reduce_scatter_with_fp32_accumulation=True" @@ -945,8 +956,9 @@ def group_params_for_buffers( Each distinct buffer is identified by a BufferKey with three dimensions: - param_dtype: storage dtype (torch.uint8 for FP8/NVFP4 parameters, else param.dtype). - grad_dtype: gradient reduction dtype (torch.float if grad_reduce_in_fp32, else param.dtype). - - is_expert_parallel: whether the parameter is expert-parallel (param.allreduce == False), - which requires a separate buffer with a different data-parallel group. + - is_expert_parallel: whether the parameter uses the expert topology (param.allreduce == False), + which requires a separate buffer for the expert data-parallel group. This is true for experts + when expert-parallelism > 1 or expert-tensor-parallelism != tensor-parallelism. The param_indices track each parameter's position among same-dtype params (using the "fake" high-precision dtype for FP8/NVFP4 params), needed for loading non-native-fp8 @@ -1156,6 +1168,7 @@ def __init__( param_layout = _compute_default_per_buffer_param_layout(self.params, bucket_size) self.param_index_map = param_layout.param_index_map self.bucket_indices = param_layout.bucket_indices + self.num_optimizer_shards = param_layout.num_optimizer_shards per_bucket_numel_unpadded = param_layout.per_bucket_numel_unpadded # Check if this buffer contains NVFP4 params. diff --git a/megatron/core/distributed/torch_fully_sharded_data_parallel.py b/megatron/core/distributed/torch_fully_sharded_data_parallel.py index 5babb6312ac..cd8aab8aba6 100644 --- a/megatron/core/distributed/torch_fully_sharded_data_parallel.py +++ b/megatron/core/distributed/torch_fully_sharded_data_parallel.py @@ -96,6 +96,8 @@ def save_custom_attrs(module): # micro-batch id, thus removing unnecessary memory stores attrs['_fp8_attrs']['transpose_invalid'] = False del attrs['_fp8_attrs']['transpose'] + # Mark this parameter as an FSDP2 parameter. + attrs["is_torch_fsdp2_param"] = True custom_attrs[name] = {k: v for k, v in attrs.items()} return custom_attrs diff --git a/megatron/core/extensions/transformer_engine.py b/megatron/core/extensions/transformer_engine.py index 8d797e816db..d9e4d4a2a27 100644 --- a/megatron/core/extensions/transformer_engine.py +++ b/megatron/core/extensions/transformer_engine.py @@ -8,8 +8,9 @@ import io import os import pickle +import re import warnings -from contextlib import nullcontext +from contextlib import contextmanager, nullcontext from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Set, Tuple, cast import torch @@ -19,7 +20,7 @@ from torch.nn.parameter import Parameter from typing_extensions import override -from megatron.core.dist_checkpointing.mapping import ShardedStateDict +from megatron.core.dist_checkpointing.mapping import ShardedObject, ShardedStateDict from megatron.core.dist_checkpointing.utils import replace_prefix_for_sharding from megatron.core.enums import Fp4Recipe, Fp8Recipe from megatron.core.extensions.transformer_engine_int4_fake_qat import ( @@ -87,6 +88,42 @@ HAVE_TE = False _TE_CONFIG_TYPE_KEY = "transformer_engine_config_type" +_EXPERT_PARAMETER_NAME_PATTERN = re.compile(r"(weight|bias)\d*") + + +def _set_expert_parameter_attributes( + module: torch.nn.Module, parallel_mode: Optional[str], use_expert_pgs: bool +) -> None: + """Set process-group and tensor-partition metadata on an expert TE module. + + ``allreduce=False`` selects EDP for gradient reduction. + + Weights and biases, including TEGroupedLinear's numbered parameters, are also marked as + TP-partitioned according to ``parallel_mode``; row-parallel biases remain replicated. + + Any parameter which is partitioned along TP or ETP is marked with ``tensor_model_parallel``, + which ensures that all shards contribute to the gradient norm. + + Args: + module: Transformer Engine module whose direct parameters should be marked. + parallel_mode: Tensor-parallel mode used by the module (``"column"``, ``"row"``, or None). + use_expert_pgs: Whether to use EP/ETP/EDP process groups instead of TP/CP/DP. + """ + for name, param in module.named_parameters(recurse=False): + param.allreduce = not use_expert_pgs + + name_match = _EXPERT_PARAMETER_NAME_PATTERN.fullmatch(name) + parameter_kind = name_match.group(1) if name_match else None + is_weight = parameter_kind == "weight" + is_bias = parameter_kind == "bias" + is_partitioned = parallel_mode in ("column", "row") and ( + is_weight or (parallel_mode == "column" and is_bias) + ) + if is_weight or is_bias: + param.tensor_model_parallel = is_partitioned + if is_partitioned: + param.partition_dim = 1 if parallel_mode == "row" else 0 + param.partition_stride = 1 class TransformerEngineConfigType(enum.Enum): @@ -384,6 +421,90 @@ def condition_init_method(config, init_method): return init_method if config.perform_initialization else (lambda w: None) +def _gtp_pre_init( + module, + output_size, + gtp_remat_group, + extra_kwargs, + *, + is_expert=False, + rng_via_kwarg=True, + out_split_size=1, +): + """Pre-shard ``out_features`` so plain TE builds this rank's shard; route init to a per-rank + RNG region (``rng_via_kwarg=False`` for LayerNormLinear). Returns ``(out_features, gtp_ctx)``. + + ``out_split_size`` is the factor TE further splits ``out_features`` by AFTER GTP (=tp_size for + column-parallel, else 1). GTP pads the per-TP slice (``output_size // out_split_size``) so each + rank's final shard stays alignment-divisible. Padding the full ``out_features`` would leave the + post-TP-split shard mis-aligned (MXFP8 needs dims divisible by 32). + """ + from megatron.core.tensor_parallel.gtp_api import gtp_remat_shard_dim0 + from megatron.core.tensor_parallel.random import get_gtp_remat_rng_tracker_name + + assert ( + output_size % out_split_size == 0 + ), f"_gtp_pre_init: output_size={output_size} not divisible by out_split_size={out_split_size}" + per_rank, pad_length = gtp_remat_shard_dim0(output_size // out_split_size, gtp_remat_group) + shard_out = per_rank * out_split_size + gtp_ctx = (gtp_remat_group, pad_length, output_size) + + tracker_name = get_gtp_remat_rng_tracker_name(is_expert=is_expert) + if rng_via_kwarg: + extra_kwargs["rng_tracker_name"] = tracker_name + else: + module.rng_tracker_name = tracker_name + return shard_out, gtp_ctx + + +def _gtp_attach_post_init(module, gtp_ctx, is_grouped=False): + """Attach the GTP surface to a pre-sharded TE module's weights and restore logical out_features. + + ``is_grouped=True`` for GroupedLinear (per-expert weight0..N, coalesced AG via weight_list). + """ + from megatron.core.tensor_parallel.gtp_api import attach_gtp_to_presharded_module + + gtp_remat_group, pad_length, logical_out_features = gtp_ctx + # Restore the LOGICAL out_features (the sharded value was only needed to size the weight in + # super().__init__): downstream code reads it, e.g. the grouped-MLP fusion gate checks + # fc1.out_features == 2 * fc2.in_features (a shard-sized fc1 would silently disable fusion). + module.out_features = logical_out_features + attach_gtp_to_presharded_module(module, gtp_remat_group, pad_length, is_grouped=is_grouped) + + +@contextmanager +def _init_gtp_remat_context( + module, + output_size, + gtp_remat_group, + extra_kwargs, + *, + is_expert=False, + is_grouped=False, + rng_via_kwarg=True, + out_split_size=1, +): + """Wrap a plain TE constructor: yield out_features for ``super().__init__`` (pre-sharded under + GTP), then attach GTP wiring on exit (skipped if construction raises, so it can't half-init). + + ``out_split_size`` = tp_size TE splits ``out_features`` by after GTP (column-parallel), else 1. + """ + if gtp_remat_group is None or gtp_remat_group.size() <= 1: + yield output_size + return + out_features, gtp_ctx = _gtp_pre_init( + module, + output_size, + gtp_remat_group, + extra_kwargs, + is_expert=is_expert, + rng_via_kwarg=rng_via_kwarg, + out_split_size=out_split_size, + ) + yield out_features + _gtp_attach_post_init(module, gtp_ctx, is_grouped=is_grouped) + + def split_te_layernorm_column_parallel_linear( fused_layer, config, @@ -765,6 +886,7 @@ def __init__( symmetric_ar_type: Optional[str] = None, tp_group: Optional[torch.distributed.ProcessGroup] = None, name: str | None = None, + gtp_remat_group: Optional[torch.distributed.ProcessGroup] = None, ): """ Args: @@ -860,6 +982,10 @@ def __init__( tp_size = get_pg_size(tp_group) self.expert_parallel = self.config.expert_model_parallel_size > 1 + use_expert_pgs = is_expert and ( + self.expert_parallel + or self.config.expert_tensor_parallel_size != self.config.tensor_model_parallel_size + ) if is_expert: rng_tracker_name = get_expert_parallel_rng_tracker_name() else: @@ -897,8 +1023,16 @@ def __init__( init_quant_context = _get_fp8_model_init_for_quant_params( self.te_quant_params, torch.is_grad_enabled() ) + init_gtp_remat_context = _init_gtp_remat_context( + self, + output_size, + gtp_remat_group, + extra_kwargs, + is_expert=is_expert, + out_split_size=tp_size if te_parallel_mode == "column" else 1, + ) - with init_quant_context: + with init_quant_context, init_gtp_remat_context as output_size: super().__init__( in_features=input_size, out_features=output_size, @@ -917,11 +1051,10 @@ def __init__( **extra_kwargs, ) - for param in self.parameters(): - if is_expert: - # Reduce the gradient on the expert_data_parallel group for expert linear layers - setattr(param, "allreduce", not self.expert_parallel) - else: + if is_expert: + _set_expert_parameter_attributes(self, parallel_mode, use_expert_pgs) + else: + for param in self.parameters(): # Reduce the gradient on DP group setattr(param, "allreduce", True) if parallel_mode == "duplicated": @@ -1104,6 +1237,10 @@ def __init__( ), "Must have at least TE version 2.3 or higher to use symmetric memory all reduce" extra_kwargs["symmetric_ar_type"] = self.config.symmetric_ar_type + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat", "expt_gtp_remat"] + ) + gtp_remat_group = pg_collection.expt_gtp_remat if is_expert else pg_collection.gtp_remat self.stride = stride self.te_quant_params: Optional[TEQuantizationParams] = None @@ -1112,11 +1249,22 @@ def __init__( init_quant_context = _get_fp8_model_init_for_quant_params( self.te_quant_params, torch.is_grad_enabled() ) + # Yield a separate gtp_output_size: the logical output_size is reused below for cpu-init + # (divide(output_size, tp_size)), so it must stay unsharded. + # rng_via_kwarg=False: TE's LayerNormLinear constructor has no rng_tracker_name kwarg. + init_gtp_remat_context = _init_gtp_remat_context( + self, + output_size, + gtp_remat_group, + extra_kwargs, + rng_via_kwarg=False, + out_split_size=self.tp_size, + ) - with init_quant_context: + with init_quant_context, init_gtp_remat_context as gtp_output_size: super().__init__( in_features=input_size, - out_features=output_size, + out_features=gtp_output_size, eps=self.config.layernorm_epsilon, sequence_parallel=self.config.sequence_parallel, fuse_wgrad_accumulation=self.config.gradient_accumulation_fusion, @@ -1220,6 +1368,11 @@ def extra_repr(self) -> str: f"out_features={self.out_features}, " f"bias={self.use_bias}, " f"TP={self.tp_size}" + + ( + f", GTP_remat={self.weight.gtp_remat_size}" + if getattr(self.weight, "gtp_remat_size", None) is not None + else "" + ) ) def backward_dw(self): @@ -1266,6 +1419,10 @@ def __init__( world_size = get_pg_size(tp_group) rank = get_pg_rank(tp_group) self.stride = stride + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat", "expt_gtp_remat"] + ) + gtp_remat_group = pg_collection.expt_gtp_remat if is_expert else pg_collection.gtp_remat super().__init__( input_size=input_size, @@ -1285,6 +1442,7 @@ def __init__( symmetric_ar_type=config.symmetric_ar_type, tp_group=tp_group, name=name, + gtp_remat_group=gtp_remat_group, ) # Set proper partition_stride @@ -1316,6 +1474,13 @@ def __init__( self.bias.zero_() setattr(self.bias, "allreduce", True) + if is_expert: + use_expert_pgs = ( + config.expert_model_parallel_size > 1 + or config.expert_tensor_parallel_size != config.tensor_model_parallel_size + ) + _set_expert_parameter_attributes(self, "column", use_expert_pgs) + def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): """Sharding along axis 0, bias sharded""" state_dict = self.state_dict(prefix="", keep_vars=True) @@ -1336,6 +1501,11 @@ def extra_repr(self) -> str: f"out_features={self.out_features}, " f"bias={self.use_bias}, " f"TP={self.tp_size}" + + ( + f", GTP_remat={self.weight.gtp_remat_size}" + if getattr(self.weight, "gtp_remat_size", None) is not None + else "" + ) ) def backward_dw(self): @@ -1508,6 +1678,10 @@ def __init__( ) tp_group = get_tensor_model_parallel_group_if_none(tp_group, is_expert=is_expert) self._tp_group = tp_group + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat", "expt_gtp_remat"] + ) + gtp_remat_group = pg_collection.expt_gtp_remat if is_expert else pg_collection.gtp_remat super().__init__( input_size=input_size, @@ -1528,6 +1702,7 @@ def __init__( symmetric_ar_type=config.symmetric_ar_type, tp_group=tp_group, name=name, + gtp_remat_group=gtp_remat_group, ) if config.use_cpu_initialization: world_size = get_pg_size(tp_group) @@ -1555,6 +1730,13 @@ def __init__( setattr(self.bias, "allreduce", True) setattr(self.bias, "sequence_parallel", config.sequence_parallel) + if is_expert: + use_expert_pgs = ( + config.expert_model_parallel_size > 1 + or config.expert_tensor_parallel_size != config.tensor_model_parallel_size + ) + _set_expert_parameter_attributes(self, "row", use_expert_pgs) + def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): """Sharding along axis 1, bias not sharded""" state_dict = self.state_dict(prefix="", keep_vars=True) @@ -1575,6 +1757,11 @@ def extra_repr(self) -> str: f"out_features={self.out_features}, " f"bias={self.use_bias}, " f"TP={self.tp_size}" + + ( + f", GTP_remat={self.weight.gtp_remat_size}" + if getattr(self.weight, "gtp_remat_size", None) is not None + else "" + ) ) def backward_dw(self): @@ -1759,6 +1946,9 @@ def __init__( self.kept_packed_seq_params.discard("tokens_per_sample") self.kept_packed_seq_params.discard("cp_partition_mode") + if get_te_version() < PkgVersion("2.2.0"): + self.kept_packed_seq_params.discard("pad_between_seqs") + if config.qk_clip or config.log_max_attention_logit: # qk-clip is only supported in TE 2.9.0 and later assert is_te_min_version("2.9.0"), "qk-clip is only supported in TE 2.9.0 and later" @@ -2006,6 +2196,10 @@ def __init__( extra_kwargs["ub_name"] = tp_comm_buffer_name self.expert_parallel = self.config.expert_model_parallel_size > 1 + use_expert_pgs = is_expert and ( + self.expert_parallel + or self.config.expert_tensor_parallel_size != self.config.tensor_model_parallel_size + ) if is_expert: extra_kwargs["rng_tracker_name"] = get_expert_parallel_rng_tracker_name() @@ -2019,6 +2213,7 @@ def __init__( self._tp_group = tp_group tp_size = get_pg_size(tp_group) tp_group_for_te = tp_group + gtp_remat_group = pg_collection.expt_gtp_remat self.explicit_expert_comm = is_expert and (tp_size > 1 or self.expert_parallel) @@ -2068,8 +2263,17 @@ def __init__( init_quant_context = _get_fp8_model_init_for_quant_params( self.te_quant_params, torch.is_grad_enabled() ) + init_gtp_remat_context = _init_gtp_remat_context( + self, + output_size, + gtp_remat_group, + extra_kwargs, + is_expert=True, + is_grouped=True, + out_split_size=tp_size if parallel_mode == "column" else 1, + ) - with init_quant_context: + with init_quant_context, init_gtp_remat_context as output_size: super().__init__( num_gemms=num_gemms, in_features=input_size, @@ -2088,23 +2292,7 @@ def __init__( **extra_kwargs, ) - for param in self.parameters(): - setattr(param, "allreduce", not (is_expert and self.expert_parallel)) - - # Explicitly stamp partition_dim and partition_stride on expert weight - # tensors when explicit_expert_comm cleared parallel_mode. TE ≤2.12 - # set these internally; TE ≥2.13 no longer does (parallel_mode=None - # is passed due to explicit_expert_comm). The resharding/refit planner - # relies on partition_dim to correctly plan TP gather/scatter operations. - # NOTE: we intentionally do NOT stamp tensor_model_parallel here — - # doing so would change num-zeros gradient counting. - if self.explicit_expert_comm and original_parallel_mode in ("column", "row"): - part_dim = 0 if original_parallel_mode == "column" else 1 - for i in range(num_gemms): - weight = getattr(self, f"weight{i}", None) - if weight is not None: - setattr(weight, "partition_dim", part_dim) - setattr(weight, "partition_stride", 1) + _set_expert_parameter_attributes(self, original_parallel_mode, use_expert_pgs) self._register_load_state_dict_pre_hook( type(self)._normalize_grouped_parameter_keys, with_module=True @@ -2340,6 +2528,8 @@ def _split_extra_state(self, state): return [state] * self.num_gemms state = self._decode_extra_state(state) + if state is None: + return [torch.empty(0, dtype=torch.uint8)] * self.num_gemms extra_states = [] extra_fp8_variables = state["extra_fp8_variables"] extra_fp8_variables["num_gemms"] = 1 @@ -2439,7 +2629,12 @@ def get_gemm_tensor(param_name: str, gemm_idx: int) -> torch.Tensor: ) if self.use_bias: sharded_state_dict[f"{prefix}bias{gemm_idx}"] = sub_sd[f"{gemm_idx}.bias"] - # Adjust replica ids - replication along DP modulo EP + # Set the expert-DP replica_id, picking the group by what EGTP_remat does to each entry: + # - _extra_state ShardedObject: REPLICATED across EGTP_remat → need distinct ids + # to avoid duplicate-writer collisions → use the full ``expt_dp_gtp_remat``. + # - weight ShardedTensor: SHARDED across EGTP_remat (distinct) → not replicas → + # elect the writer over the replicate group ``expt_dp``. + # EGTP_remat=1: the two groups coincide, so this is a no-op. for k, sh_ten in sharded_state_dict.items(): replica_id = sh_ten.replica_id assert ( @@ -2447,6 +2642,8 @@ def get_gemm_tensor(param_name: str, gemm_idx: int) -> torch.Tensor: ), f"Expected replica_id for {k} to be in (PP, TP, DP) format, got: {replica_id}" if getattr(sh_ten, "is_data_parallel_fully_shard", False): edp_replica_id = 0 + elif isinstance(sh_ten, ShardedObject): + edp_replica_id = get_pg_rank(self._pg_collection.expt_dp_gtp_remat) else: edp_replica_id = get_pg_rank(self._pg_collection.expt_dp) sh_ten.replica_id = (*replica_id[:2], edp_replica_id) @@ -2460,6 +2657,17 @@ def backward_dw(self): if self.delay_wgrad_compute: super().backward_dw() + def __repr__(self): + gtp_remat = getattr(getattr(self, "weight0", None), "gtp_remat_size", None) + gtp_str = f", GTP_remat={gtp_remat}" if gtp_remat is not None else "" + return ( + f"{type(self).__name__}(per expert([" + f"in={self.in_features}, out={self.out_features}]) " + f"X num_gemms={self.num_gemms}, " + f"bias={self.use_bias}, TP={self.tp_size}" + f"{gtp_str})" + ) + class TEColumnParallelGroupedLinear(TEGroupedLinear): """ Wrapper for the Transformer-Engine's `GroupedLinear` layer but specialized diff --git a/megatron/core/fp8_utils.py b/megatron/core/fp8_utils.py index ac36d40d848..4fadbd25c97 100644 --- a/megatron/core/fp8_utils.py +++ b/megatron/core/fp8_utils.py @@ -226,6 +226,46 @@ def copy_tensor_to_quantized_param(param: torch.Tensor, src: torch.Tensor) -> No dst.copy_(src.view(dst.shape)) +def copy_tensors_to_quantized_params(params: List[torch.Tensor], srcs: List[torch.Tensor]) -> None: + """List form of :func:`copy_tensor_to_quantized_param`, for a whole bucket of params. + + Same values, minus the per-param ``copy_`` and tensor-subclass dispatch: the quantizer is + resolved up front and called directly. Cast kernels are unchanged, one per param. Worth it + because those casts are small and issuing them is expensive, and under + --reuse-grad-buf-for-mxfp8-param-ag they run inside the forward pass. + + Args: + params: quantized model params to write into. + srcs: high-precision source values, one per param, in the same order. + """ + if len(params) == 0: + return + + srcs_to_cast = [] + dsts_to_cast = [] + quantizers = [] + for param, src in zip(params, srcs): + dst = _unwrap_parameter_data(param) + quantizer = ( + None + if is_grouped_tensor_with_quantized_storage(dst) + else getattr(dst, "_quantizer", None) + ) + if quantizer is None: + # Grouped storage quantizes per member; a missing quantizer has to be built. Both + # cases are handled by the single-param path. + copy_tensor_to_quantized_param(param, src) + continue + srcs_to_cast.append(src.view(dst.shape)) + dsts_to_cast.append(dst) + quantizers.append(quantizer) + + # Equivalent to dst.copy_(src), but entered directly instead of via the aten::copy_ op, + # QuantizedTensor.__torch_dispatch__ (type and usage checks) and dst.quantize_(src). + for src, quantizer, dst in zip(srcs_to_cast, quantizers, dsts_to_cast): + quantizer.update_quantized(src, dst) + + def modify_grouped_tensor_rowwise_storage(tensor: torch.Tensor, new_storage: torch.Tensor) -> None: """Replace a high-precision Transformer Engine GroupedTensor's rowwise storage.""" tensor = _unwrap_parameter_data(tensor) diff --git a/megatron/core/inference/apis/_llm_base.py b/megatron/core/inference/apis/_llm_base.py index 93b1bda30c8..acc2ca336e3 100644 --- a/megatron/core/inference/apis/_llm_base.py +++ b/megatron/core/inference/apis/_llm_base.py @@ -6,7 +6,8 @@ ``MegatronAsyncLLM``: ``_EventLoopManager``, ``_CoordinatorRuntime``, and ``_MegatronLLMBase``. The public sync/async wrappers live on the subclasses; this base only exposes shared engine state, runtime spawn, validation -helpers, and the private ``__impl`` coroutines. +helpers, the public sync bridge (``submit``/``run_sync``), and the private +``__impl`` coroutines. """ import asyncio @@ -296,6 +297,7 @@ def __init__( self._loop_manager: "Optional[_EventLoopManager]" = None self._coord_runtime: "Optional[_CoordinatorRuntime]" = None self._shutdown_called: bool = False + self._serve_started: bool = False if use_coordinator: loop_manager = _EventLoopManager() @@ -342,8 +344,61 @@ def controller(self) -> "TextGenerationController": """The underlying :class:`TextGenerationController`.""" return self._controller + # ---- sync bridge (public) ---- + + def submit(self, coro: Coroutine) -> "concurrent.futures.Future": + """Schedule ``coro`` on the background runtime loop; return its future. + + The returned :class:`concurrent.futures.Future` can be consumed from + any context: block with ``.result()`` from sync code, or wrap with + ``asyncio.wrap_future(...)`` and ``await`` it from a coroutine. + Callable from any thread, including threads whose own event loop is + running (e.g. an embedder's dispatch loop) -- the coroutine executes + on the runtime loop either way. + + Raises: + RuntimeError: in direct mode (``use_coordinator=False``), which + has no background runtime loop. + """ + self._assert_coordinator() + assert self._loop_manager is not None + return self._loop_manager.submit(coro) + + def run_sync(self, coro: Coroutine): + """Schedule ``coro`` on the background runtime loop and block on it. + + Safe to call from any thread except the runtime loop itself (that + would deadlock and raises instead). Calling from a thread whose own + event loop is running is allowed: the caller's loop stalls until the + result returns, while ``coro`` runs on the runtime loop. + + Raises: + RuntimeError: in direct mode (``use_coordinator=False``), or when + called from a coroutine running on the runtime loop itself. + """ + self._assert_coordinator() + assert self._loop_manager is not None + return self._loop_manager.run_sync(coro) + # ---- internal helpers ---- + def _stop_frontend_if_started(self) -> None: + """Stop the HTTP frontend if ``serve()`` started one on this rank. + + Called first by both facades' ``shutdown()`` so no new requests + arrive while the coordinator is torn down. Invariant: + ``_serve_started`` can only be True when ``use_coordinator=True`` + because ``serve()`` raises otherwise. + """ + if not self._serve_started: + return + from megatron.core.inference.text_generation_server.dynamic_text_gen_server.text_generation_server import ( # pylint: disable=line-too-long + stop_text_gen_server, + ) + + stop_text_gen_server() + self._serve_started = False + def _assert_primary(self) -> None: if not self._is_primary_rank: raise RuntimeError( diff --git a/megatron/core/inference/apis/async_llm.py b/megatron/core/inference/apis/async_llm.py index a64fd07a78a..2c6b20693f3 100644 --- a/megatron/core/inference/apis/async_llm.py +++ b/megatron/core/inference/apis/async_llm.py @@ -61,8 +61,6 @@ def __init__( coordinator_host=coordinator_host, coordinator_port=coordinator_port, ) - # Set in serve() when this rank starts the HTTP frontend; consulted by shutdown(). - self._serve_started: bool = False async def generate( self, @@ -147,17 +145,7 @@ async def shutdown(self) -> None: return self._shutdown_called = True - # If we started an HTTP frontend, stop it first so no new requests - # arrive while we tear down the coordinator. Invariant: - # ``_serve_started`` can only be True when ``use_coordinator=True`` - # because ``serve()`` raises otherwise. - if self._serve_started: - from megatron.core.inference.text_generation_server.dynamic_text_gen_server.text_generation_server import ( # pylint: disable=line-too-long - stop_text_gen_server, - ) - - stop_text_gen_server() - self._serve_started = False + self._stop_frontend_if_started() if not self._use_coordinator: return diff --git a/megatron/core/inference/apis/llm.py b/megatron/core/inference/apis/llm.py index 40222948987..6bbdce570d8 100644 --- a/megatron/core/inference/apis/llm.py +++ b/megatron/core/inference/apis/llm.py @@ -5,6 +5,7 @@ from typing import List, Optional, Union from megatron.core.inference.apis._llm_base import _MegatronLLMBase +from megatron.core.inference.apis.serve_config import ServeConfig from megatron.core.inference.config import InferenceConfig from megatron.core.inference.inference_request import DynamicInferenceRequest from megatron.core.inference.sampling_params import SamplingParams @@ -24,12 +25,9 @@ class MegatronLLM(_MegatronLLMBase): - Sync lifecycle controls: :meth:`pause` / :meth:`unpause` / :meth:`suspend` / :meth:`resume` / :meth:`shutdown` / :meth:`wait_for_shutdown`. + - :meth:`serve` for OpenAI-compatible HTTP serving on the primary rank. - Context-manager protocol: ``with MegatronLLM(...) as llm:``; exit calls :meth:`shutdown`. - - Note: - ``serve()`` (online HTTP serving) is async-only by design; use - :class:`MegatronAsyncLLM` for serving. """ def __init__( @@ -132,6 +130,7 @@ def shutdown(self) -> None: if self._shutdown_called: return self._shutdown_called = True + self._stop_frontend_if_started() if not self._use_coordinator: return # direct mode: nothing to tear down assert self._loop_manager is not None @@ -139,6 +138,55 @@ def shutdown(self) -> None: # Sync caller already on its own thread; no need for to_thread. self._loop_manager.stop() + def serve(self, serve_config: ServeConfig, *, blocking: bool = True) -> None: + """Start the OpenAI-compatible HTTP frontend. + + Coordinator mode only. The HTTP frontend runs only on the primary + rank (global rank 0); other ranks no-op the HTTP setup but still + respect ``blocking`` (so all ranks return together). + + With ``blocking=True`` (default), this blocks the calling thread until + the engine loop terminates via :meth:`shutdown` -- suitable for + standalone serving scripts. With ``blocking=False``, this returns once + the HTTP frontend is up (primary) or immediately (workers); the engine + loop continues in the background runtime, and the user can call + :meth:`generate` / :meth:`shutdown` afterward. + + Raises: + ValueError: if ``use_coordinator=False`` (HTTP serving requires + the coordinator path). + """ + if not self._use_coordinator: + raise ValueError("MegatronLLM.serve() requires use_coordinator=True") + + if self._is_primary_rank: + # Lazy import: keep the module importable in environments where + # the HTTP server backend (Quart/Hypercorn) isn't installed. + import torch.distributed as dist + + from megatron.core.inference.text_generation_server.dynamic_text_gen_server.text_generation_server import ( # pylint: disable=line-too-long + start_text_gen_server, + ) + + assert self._coord_runtime is not None + start_text_gen_server( + coordinator_addr=self._coord_runtime.coord_addr, + tokenizer=self._controller.tokenizer, + rank=dist.get_rank(), + server_port=serve_config.port, + parsers=serve_config.parsers, + verbose=serve_config.verbose, + num_replicas=serve_config.frontend_replicas, + hostname=serve_config.host, + ) + self._serve_started = True + + if blocking: + # Block until the engine loop terminates (shutdown was invoked + # somewhere in this process; for serve(blocking=True) typically by + # SIGINT or out-of-band orchestration). + self.wait_for_shutdown() + def wait_for_shutdown(self) -> None: """Block until the engine loop terminates. Direct mode no-op.""" if not self._use_coordinator: diff --git a/megatron/core/inference/config.py b/megatron/core/inference/config.py index 1d9541207e7..46d2dbca1b0 100644 --- a/megatron/core/inference/config.py +++ b/megatron/core/inference/config.py @@ -1,5 +1,6 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +import warnings from dataclasses import InitVar, dataclass from enum import Enum from typing import List, Literal, Optional, Tuple @@ -100,8 +101,8 @@ class PrefixCachingCoordinatorPolicy(str, Enum): FIRST_PREFIX_BLOCK = "first_prefix_block" """Route to the rank that has the first block hash cached. O(ranks) check.""" - ROUND_ROBIN = "round_robin" - """Route requests to ranks in round-robin order, ignoring prefix affinity.""" + LOAD_BALANCED = "load_balanced" + """Route to the rank with the fewest in-flight requests. Ignores prefix affinity.""" class KVCacheManagementMode(str, Enum): @@ -139,8 +140,8 @@ class AsyncScheduleMode(str, Enum): LEGACY = "legacy" """Resolve requests before preparing the next forward pass.""" - SERIAL = "serial" - """Prepare and forward speculatively before resolving the sampled requests.""" + ASYNC = "async" + """Overlap asynchronous scheduling phases by reordering them to prepare-before-resolve.""" @dataclass @@ -217,7 +218,7 @@ class InferenceConfig: Maximum number of cuda graphs to capture. Graph token counts are spaced from 1 up to a per-graph-type budget: - Decode-only graphs are always bounded by `max_requests * (num_speculative_tokens + 1)`. - - Prefill/mixed graphs share that same bound by default, + - Prefill/mixed graphs are bounded by `cuda_graph_max_tokens` by default, or extend up to `max_tokens` when `cuda_graph_all_prefills` is set. Due to rounding, the actual number of cuda graphs may not equal this argument. """ @@ -245,11 +246,19 @@ class InferenceConfig: cuda_graph_all_prefills: bool = False """ Whether prefill/mixed CUDA graphs should span up to `max_tokens`. - When False (default), prefill/mixed graphs are bounded by the same token limit as decode graphs: - `max_requests * (num_speculative_tokens + 1)`. + When False (default), prefill/mixed graphs are bounded by `cuda_graph_max_tokens`. When True, prefill/mixed graph capture is extended to cover the full `max_tokens` budget. """ + cuda_graph_max_tokens: int = 512 + """ + Token ceiling for the largest captured prefill/mixed CUDA graph. + This is a raw token count (not scaled by speculative decoding). The effective ceiling is + clamped to `[max_requests * (num_speculative_tokens + 1), max_tokens]` so it never falls + below the decode bound nor exceeds the token budget. Ignored when `cuda_graph_all_prefills` + is set, which extends capture to the full `max_tokens`. + """ + static_kv_memory_pointers: bool = False """ Whether the KV cache (and Mamba states) will reside at the same memory addresses @@ -300,7 +309,7 @@ class InferenceConfig: """ prefix_caching_coordinator_policy: PrefixCachingCoordinatorPolicy = ( - PrefixCachingCoordinatorPolicy.FIRST_PREFIX_BLOCK + PrefixCachingCoordinatorPolicy.LOAD_BALANCED ) """Routing policy for the DP inference coordinator. See `PrefixCachingCoordinatorPolicy` for options. @@ -318,7 +327,17 @@ class InferenceConfig: """GPU memory budget (in GB) for the Mamba state cache used by prefix caching on hybrid models. Each cache slot stores SSM and conv states for all Mamba layers at a single block boundary. When set, Mamba states at KV divergence and last-aligned - block boundaries are cached and reused across requests with matching prefixes.""" + block boundaries are cached and reused across requests with matching prefixes. + + This budget covers both buffers allocated by MambaSlotAllocator: the durable cache + (ssm_states/conv_states, max_slots slots reused across requests) and the per-step + extraction scratch (intermediate_ssm_out/intermediate_conv_out). The scratch is + sized to the tighter of two per-step bounds, + ``min(ceil(max_tokens / block_size_tokens), 3 * max_requests)``, since a single + engine step can extract at most one state per block_size_tokens of its token budget + (and at most 3 per request). The scratch is reserved from this budget first, so a + smaller ``max_tokens`` (or ``max_requests``) shrinks the scratch and leaves more + durable cache slots.""" # ================================= # Logging config @@ -346,7 +365,17 @@ class InferenceConfig: """ sampling_backend: Literal['torch', 'flashinfer'] = 'torch' - """Which sampling kernels to use during inference.""" + """Which sampling kernels to use during inference. Falls back to "torch" with a warning if + "flashinfer" is requested but the package is not installed.""" + + offset_sampling_seed_by_dp_rank: bool = True + """ + If True, offset `inference_sampling_seed` by the data-parallel rank when seeding the + sampling RNG. This gives each DP rank a unique generation seed so that the same prompt + routed to different ranks produces different samples (important for RL training). + If False (or `ModelParallelConfig.deterministic_mode` / `--deterministic-mode` is + enabled), then all DP ranks share the same sampling / generation seed. + """ async_sched_mode: AsyncScheduleMode = AsyncScheduleMode.LEGACY """Mode used to schedule dynamic batching inference work.""" @@ -414,8 +443,9 @@ def __post_init__(self, verbose: bool): if self.sampling_backend == 'flashinfer': try: import flashinfer # noqa: F401 - except ImportError as e: - raise ImportError( - "sampling_backend='flashinfer' requires the flashinfer package; " - "install it or set sampling_backend='torch'." - ) from e + except ImportError: + warnings.warn( + "sampling_backend='flashinfer' was requested but the flashinfer " + "package is not installed; falling back to sampling_backend='torch'." + ) + self.sampling_backend = 'torch' diff --git a/megatron/core/inference/contexts/attention_context/mamba_metadata.py b/megatron/core/inference/contexts/attention_context/mamba_metadata.py index 3e98f0324e6..9984b2dd71a 100644 --- a/megatron/core/inference/contexts/attention_context/mamba_metadata.py +++ b/megatron/core/inference/contexts/attention_context/mamba_metadata.py @@ -14,7 +14,13 @@ class MambaMetadata: """Manages the metadata tensors required for Mamba layers during inference.""" def __init__( - self, max_requests: int, max_tokens: int, mamba_chunk_size: int = 128, d_conv: int = 0 + self, + max_requests: int, + max_tokens: int, + *, + max_intermediate_count: int, + mamba_chunk_size: int = 128, + d_conv: int = 0, ): """ Initializes the Mamba slot allocator. @@ -22,6 +28,11 @@ def __init__( Args: max_requests (int): The maximum number of concurrent requests. max_tokens (int): The maximum number of tokens. + max_intermediate_count (int): Per-step upper bound on Mamba + intermediate-state extractions; sizes the intermediate metadata + buffers. Computed once by DynamicInferenceContext (as + max_mamba_intermediate_states_per_step) and shared with + MambaSlotAllocator. mamba_chunk_size (int): The chunk size used by the Mamba SSM Triton kernels. d_conv (int): Convolution window size (from mamba_conv_states_shape[-1]). Used for vectorized conv state extraction at intermediate offsets. @@ -91,22 +102,20 @@ def __init__( ) self.mamba_state_free_slot_count = self.max_requests - # Intermediate state extraction buffers (CUDA graph compatible) - # Each prefill request can produce up to 3 intermediate offsets - self.max_intermediate_count = MAX_INTERMEDIATE_OFFSETS_PER_REQUEST * max_requests + # Intermediate state extraction buffers (CUDA graph compatible). Sized by + # the per-step token-budget cap shared from DynamicInferenceContext. + self.max_intermediate_count = max_intermediate_count self._intermediate_chunk_indices_buffer = torch.zeros( self.max_intermediate_count, dtype=torch.int64, device=self.device ) self._intermediate_abs_positions_buffer = torch.full( (self.max_intermediate_count,), d_conv, dtype=torch.int32, device=self.device ) - # Constant gather offsets for conv state extraction: [-d_conv, ..., -1] - if d_conv > 0: - self.conv_gather_offsets = torch.arange( - -d_conv, 0, dtype=torch.int32, device=self.device - ) - else: - self.conv_gather_offsets = None + # Runtime real-count tensor read by the fused gather+scatter Triton + # kernels (intermediate_extraction.py). Fixed-address, rewritten each step + # so captured CUDA graphs stay valid while the kernels skip padded slots + # (pid_slot >= real_count). + self._intermediate_real_count_buffer = torch.zeros(1, dtype=torch.int32, device=self.device) # Coalesced production path: pinned CPU views + shared GPU views bound # by DynamicInferenceContext so that the per-step Mamba metadata fields @@ -169,6 +178,7 @@ def reset_varlen_metadata(self) -> None: # Intermediate state extraction views self.intermediate_chunk_indices = None self.intermediate_abs_positions = None + self.intermediate_real_count = None self.intermediate_count = 0 self.per_request_intermediate_counts = [] @@ -381,13 +391,24 @@ def _update_intermediate_metadata( intermediate_counts_gpu: [real_prefill_count] int32 GPU tensor of per-request offset counts (0-3), or None. real_prefill_count: Number of real (non-padding) prefill requests. + padded_prefill_count: Prefill request count after batch padding + (equals the captured graph bucket under CUDA graphs, or the + round-up-padded count in eager mode; always >= real_prefill_count). + Bounds the exposed/padded extent of the intermediate views via + ``max_count`` so CUDA graph replay always touches a fixed-size + region within the scratch buffers. cu_seqlens_gpu: GPU cu_seqlens tensor to read from. Defaults to the legacy standalone ``_cu_seqlens_buffer`` used by :meth:`update`; the coalesced production path passes the shared ``ContextGPUView.mamba_cu_seqlens`` view. """ chunk_size = self.mamba_chunk_size - max_count = padded_prefill_count * MAX_INTERMEDIATE_OFFSETS_PER_REQUEST + # Cap at the token-budget bound so the per-step views never exceed the + # buffers, even for high-prefill-count graph buckets where + # padded_prefill_count * MAX_INTERMEDIATE_OFFSETS_PER_REQUEST would. + max_count = min( + padded_prefill_count * MAX_INTERMEDIATE_OFFSETS_PER_REQUEST, self.max_intermediate_count + ) if cu_seqlens_gpu is None: cu_seqlens_gpu = self._cu_seqlens_buffer @@ -438,6 +459,13 @@ def _update_intermediate_metadata( valid_abs_positions = abs_positions_2d[valid_mask] real_count = valid_chunk_indices.numel() + # The token-budget bound guarantees this; fail loudly rather than + # silently overrun the scratch buffers if the candidate-offset + # logic in MambaSlotAllocator.compute_and_store_offsets changes. + assert real_count <= self.max_intermediate_count, ( + f"Mamba intermediate count {real_count} exceeds buffer size " + f"{self.max_intermediate_count}" + ) self._intermediate_chunk_indices_buffer[:real_count] = valid_chunk_indices self._intermediate_abs_positions_buffer[:real_count] = valid_abs_positions.to( torch.int32 @@ -445,8 +473,11 @@ def _update_intermediate_metadata( # Pad unused slots with safe defaults for CUDA graph replay: # - chunk_indices=0: reads from chunk 0 (always exists), output ignored - # - abs_positions=d_conv: conv gather reads tokens [0..d_conv-1], - # which are within bounds and produce a valid but unused state + # - abs_positions=d_conv: conv gather reads tokens [0..d_conv-1]. + # These are within bounds only when the prefill has at least + # d_conv tokens; shorter sequences (e.g. small CUDA-graph warmup + # buckets) would overrun the token axis, so _ssm_prefill clamps + # the gather positions into range. The gathered state is unused. if real_count < max_count: self._intermediate_chunk_indices_buffer[real_count:max_count].fill_(0) self._intermediate_abs_positions_buffer[real_count:max_count].fill_(self.d_conv) @@ -462,15 +493,24 @@ def _update_intermediate_metadata( self.intermediate_chunk_indices = self._intermediate_chunk_indices_buffer[:max_count] self.intermediate_abs_positions = self._intermediate_abs_positions_buffer[:max_count] + # Publish real_count to the fixed-address GPU tensor the scatter + # kernels consult. fill_ is async (no host sync) and keeps the tensor + # at the same address captured graphs reference. + self._intermediate_real_count_buffer.fill_(self.intermediate_count) + self.intermediate_real_count = self._intermediate_real_count_buffer else: # No extraction: fill with safe defaults for CUDA graph warmup - # (same rationale as padding comment above) + # (same rationale as padding comment above; abs_positions=d_conv may + # exceed a sub-d_conv warmup sequence, so _ssm_prefill clamps the + # gather positions into range and the gathered state is unused) self._intermediate_chunk_indices_buffer[:max_count] = 0 self._intermediate_abs_positions_buffer[:max_count] = self.d_conv self.intermediate_count = 0 self.per_request_intermediate_counts = [] self.intermediate_chunk_indices = self._intermediate_chunk_indices_buffer[:max_count] self.intermediate_abs_positions = self._intermediate_abs_positions_buffer[:max_count] + self._intermediate_real_count_buffer.fill_(0) + self.intermediate_real_count = self._intermediate_real_count_buffer def compute_cpu_metadata( self, diff --git a/megatron/core/inference/contexts/dynamic_context.py b/megatron/core/inference/contexts/dynamic_context.py index 4b3e36031b7..b938b4f2c5e 100644 --- a/megatron/core/inference/contexts/dynamic_context.py +++ b/megatron/core/inference/contexts/dynamic_context.py @@ -36,7 +36,9 @@ ) from megatron.core.package_info import __version__ as mcore_version from megatron.core.transformer import MLATransformerConfig, TransformerConfig +from megatron.core.transformer.enums import InferenceCudaGraphScope from megatron.core.transformer.moe.token_dispatcher_inference import ( + InferenceAllGatherDispatcherBase, NCCLAllGatherDispatcher, NVLSAllGatherVDispatcher, ) @@ -49,7 +51,7 @@ from .base_context import BaseInferenceContext from .gpu_view import ContextGPUView from .kv_block_allocator import KVBlockAllocator -from .mamba_slot_allocator import MambaSlotAllocator +from .mamba_slot_allocator import MAX_INTERMEDIATE_OFFSETS_PER_REQUEST, MambaSlotAllocator from .routing_metadata import RoutingMetadata try: @@ -276,6 +278,13 @@ def __init__(self, model_config: TransformerConfig, inference_config: InferenceC # Prefix caching hit tracking (accumulated, reset by engine after logging). self.prefix_cache_hits = 0 # requests that matched at least one cached block self.prefix_cache_blocks_matched = 0 # total matched blocks across all requests + # Prefill compute accounting (drained into engine accumulators each step). + # computed = prompt tokens actually run through the model this step; + # skipped = prompt tokens whose prefill was skipped via a prefix-cache hit. + # A high skipped fraction confirms prefix caching is saving prefill compute + # (so any per-step latency growth is attention-over-context, not re-prefill). + self.prefix_cache_prefill_computed_tokens = 0 + self.prefix_cache_prefill_skipped_tokens = 0 # Engine step counter (used for logging, metrics, and event tracking) self.step_count = 0 @@ -311,6 +320,7 @@ def __init__(self, model_config: TransformerConfig, inference_config: InferenceC f"num_speculative_tokens ({self.num_speculative_tokens}) must be < " f"block_size_tokens ({inference_config.block_size_tokens})" ) + self._async_sched_token_offsets = None # Cache the PP group we should use for PP collectives inside the context. # If the model provides a pg_collection with a pp group, prefer it. @@ -557,6 +567,8 @@ def __init__(self, model_config: TransformerConfig, inference_config: InferenceC # Initialize context state. self.params_dtype = model_config.params_dtype + self.hidden_size = model_config.hidden_size + self.inference_cuda_graph_scope = model_config.inference_cuda_graph_scope self.max_sequence_length = inference_config.max_sequence_length # Block ids. With speculative decoding, blocks are pre-allocated when the @@ -587,6 +599,15 @@ def __init__(self, model_config: TransformerConfig, inference_config: InferenceC self.max_tokens = inference_config.max_tokens or self.DEFAULT_MAX_TOKENS + # Per-step upper bound on Mamba intermediate-state extractions, shared with + # MambaMetadata and MambaSlotAllocator so scratch/metadata buffers and the + # budget accounting agree. Bounded both by the token budget (one block + # boundary per block_size_tokens) and by the request budget + # (MAX_INTERMEDIATE_OFFSETS_PER_REQUEST per request); + token_based_count = math.ceil(self.max_tokens / self.block_size_tokens) + request_based_count = MAX_INTERMEDIATE_OFFSETS_PER_REQUEST * self.max_requests + self.max_mamba_intermediate_states_per_step = min(token_based_count, request_based_count) + assert self.max_tokens >= self.max_requests, ( f"max_tokens ({self.max_tokens}) must be >= " f"max_requests ({self.max_requests}), " @@ -627,6 +648,12 @@ def __init__(self, model_config: TransformerConfig, inference_config: InferenceC and model_config.inference_moe_token_dispatcher_type == 'nccl' ) + # are we using the inference_optimized nvls ep dispatcher for MoEs? + self._nvls_dispatcher = ( + get_pg_size(self.expert_model_parallel_group) > 1 + and model_config.inference_moe_token_dispatcher_type == 'nvls' + ) + # are we using the training a2a dispatcher for MoEs? # Note that this is not optimal for speed. self._training_ep_dispatcher = ( @@ -645,12 +672,15 @@ def __init__(self, model_config: TransformerConfig, inference_config: InferenceC ) # CUDA graph token budget for prefill/mixed graphs. Decode graphs are always - # capped at max_requests * (num_speculative_tokens + 1) inside the helper; this - # only widens the prefill/mixed range when `cuda_graph_all_prefills` is set. + # capped at max_requests * (num_speculative_tokens + 1) inside the helper. By + # default the prefill/mixed range is bounded by `cuda_graph_max_tokens`, clamped + # to never fall below that decode bound nor exceed `max_tokens`; setting + # `cuda_graph_all_prefills` widens the range to the full `max_tokens`. + decode_bound = self.max_requests * (self.num_speculative_tokens + 1) cuda_graph_max_tokens = ( self.max_tokens if inference_config.cuda_graph_all_prefills - else self.max_requests * (self.num_speculative_tokens + 1) + else min(max(inference_config.cuda_graph_max_tokens, decode_bound), self.max_tokens) ) # CUDA graph config list. @@ -671,19 +701,24 @@ def __init__(self, model_config: TransformerConfig, inference_config: InferenceC # Allocate per-step dispatcher buffers upfront so update_metadata never # triggers an allocation inside a captured CUDA graph. - if get_pg_size(self.expert_model_parallel_group) > 1: - if self._nccl_ep_dispatcher: - NCCLAllGatherDispatcher.allocate_buffers() - else: - # Use moe_latent_size if set (latent MoE: SuperV3, UltraV3), else hidden_size. - moe_hidden_size = model_config.moe_latent_size or model_config.hidden_size - NVLSAllGatherVDispatcher.allocate_buffers( - per_rank_worst_case_token_count=self.round_up_tokens(self.max_tokens) - // tp_size, - topk=model_config.moe_router_topk, - hidden_size=moe_hidden_size, - ep_group=self.expert_model_parallel_group, - ) + # + # The shared _valid_tokens_tensor scalar is read as a pointer by both fused + # MoE backends (mcore_fused_moe and vllm_fused_moe) regardless of EP size, so + # allocate it unconditionally (covers EP=1, where no dispatcher comm buffers + # exist). The EP>1 dispatchers below reallocate it as part of their own buffer + # setup, which is harmless. + InferenceAllGatherDispatcherBase.allocate_valid_tokens_tensor() + if self._nccl_ep_dispatcher: + NCCLAllGatherDispatcher.allocate_buffers() + elif self._nvls_dispatcher: + # Use moe_latent_size if set (latent MoE: SuperV3, UltraV3), else hidden_size. + moe_hidden_size = model_config.moe_latent_size or model_config.hidden_size + NVLSAllGatherVDispatcher.allocate_buffers( + per_rank_worst_case_token_count=self.round_up_tokens(self.max_tokens) // tp_size, + topk=model_config.moe_router_topk, + hidden_size=moe_hidden_size, + ep_group=self.expert_model_parallel_group, + ) # Deal with chunked prefill self.enable_chunked_prefill = inference_config.enable_chunked_prefill @@ -696,10 +731,23 @@ def __init__(self, model_config: TransformerConfig, inference_config: InferenceC self.use_flashinfer_fused_rope = inference_config.use_flashinfer_fused_rope self.inference_grouped_gemm_backend = model_config.inference_grouped_gemm_backend + # Placeholder for the MTP decoder hidden-states buffer; allocated inside + # initialize_all_tensors() when num_speculative_tokens > 0. + self.mtp_decoder_hidden_states = None + # Allocate GPU state. self.is_tensor_state_allocated = False + self._bookkeeping_no_real_work = False self.initialize_all_tensors() + # Bind the GPU real-token-count tensor onto the NVLS dispatcher class + # so it can mask out CUDA-graph padding tokens during routing. The + # tensor lives inside gpu_view._buf (fixed address) and is refreshed + # each step by transfer_bookkeeping_to_gpu(). NVLS-only — the NCCL + # dispatcher requires equal token counts across ranks already. + if self._nvls_dispatcher: + NVLSAllGatherVDispatcher.set_real_token_count_tensor(self.gpu_view.real_token_count) + # Print info. active_blocks = self.kv_block_allocator.active_count total_blocks = self.kv_block_allocator.total_count @@ -761,11 +809,23 @@ def __init__(self, model_config: TransformerConfig, inference_config: InferenceC and prefix_caching_mamba_gb > 0 ): prefix_cache_bytes = int(prefix_caching_mamba_gb * 1024**3) - prefix_cache_slots = prefix_cache_bytes // mamba_bytes_per_req + # Mirror the split done in _allocate_mamba_cache so this preview + # matches what is actually allocated: the "scratch" buffers + # (intermediate_ssm_out/intermediate_conv_out) are reserved from the + # budget first, then the rest sizes the "durable" cache + # (ssm_states/conv_states). mamba_bytes_per_req is the shared + # per-slot footprint of both. + scratch_slots = self.max_mamba_intermediate_states_per_step + scratch_bytes = scratch_slots * mamba_bytes_per_req + durable_slots = (prefix_cache_bytes - scratch_bytes) // mamba_bytes_per_req + durable_slots = max(durable_slots, 0) log_lines += [ f" Mamba prefix cache:", f" budget: {get_mem_size_str(prefix_cache_bytes)}", - f" slots: {prefix_cache_slots}", + f" extraction_scratch: {scratch_slots} slots " + f"({get_mem_size_str(scratch_bytes)})", + f" durable_slots: {durable_slots} " + f"({get_mem_size_str(durable_slots * mamba_bytes_per_req)})", f" per_slot: {get_mem_size_str(mamba_bytes_per_req)}", ] @@ -814,6 +874,7 @@ def _allocate_mamba_states(self): self.mamba_metadata = MambaMetadata( max_requests=self.max_requests, max_tokens=self.max_tokens, + max_intermediate_count=self.max_mamba_intermediate_states_per_step, mamba_chunk_size=self.mamba_chunk_size, d_conv=self.mamba_conv_states_shape[-1], ) @@ -946,6 +1007,9 @@ def initialize_all_tensors(self) -> None: device='cpu', pin_memory=True, ) + self._async_sched_token_offsets = torch.arange( + self.num_speculative_tokens + 1, device='cpu' + ) # Track request metadata. Backed by pinned CPU memory: bookkeeping is # CPU-resident; GPU consumers read from the active-slice mirror in @@ -993,6 +1057,10 @@ def initialize_all_tensors(self) -> None: _tok_int32_bytes = self.max_tokens * 4 # Request-level fields are all 4 bytes wide (5 int32 + 2 float32 = 7 fields). _req_4byte_bytes = self.max_requests * 4 + # Scalar: real (unpadded) token count for the current step. Refreshed + # in transfer_bookkeeping_to_gpu(); read on GPU via + # `gpu_view.real_token_count` (MoE routing masks padding tokens). + _real_token_count_bytes = 4 # MHA section: 5 fields (int32) shared between GraphedMHAMetadata and # NonGraphedMHAMetadata. max_bs == max_requests. _mha_query_lengths_bytes = self.max_requests * 4 @@ -1004,6 +1072,7 @@ def initialize_all_tensors(self) -> None: 3 * _tok_int64_bytes + 3 * _tok_int32_bytes + 7 * _req_4byte_bytes + + _real_token_count_bytes + _mha_query_lengths_bytes + _mha_cu_query_seq_lengths_bytes + _mha_kv_seq_lengths_bytes @@ -1130,6 +1199,14 @@ def initialize_all_tensors(self) -> None: ].view(torch.int32) _off += _req_4byte_bytes + # Scalar staging slot for the real (unpadded) token count. Refreshed + # from `self.batch_dimensions.token_count` in transfer_bookkeeping_to_gpu() + # and read on GPU via `gpu_view.real_token_count`. + self._staging_real_token_count = self._cpu_bookkeeping_buf[ + _off : _off + _real_token_count_bytes + ].view(torch.int32) + _off += _real_token_count_bytes + # Static tensor addresses to make `last_token_logits` graphable with speculative decoding. max_logit_idxs = self.max_requests * (self.num_speculative_tokens + 1) self.active_logit_idxs = torch.zeros( @@ -1221,6 +1298,7 @@ def initialize_all_tensors(self) -> None: device=torch.cuda.current_device(), max_mamba_chunks=self._max_mamba_chunks, ) + self._bookkeeping_h2d_done_event = torch.cuda.Event() # Cache of (input_ids_view, pos_ids_view) keyed by num_tokens. Instead of slicing and # unsqueezing on every new inference step (constructing new TensorImpls at 30-60 us), @@ -1267,6 +1345,40 @@ def initialize_all_tensors(self) -> None: and self.config.enable_prefix_caching ): self._allocate_mamba_cache(self.config.prefix_caching_mamba_gb) + elif self.is_hybrid_model and self.config.enable_prefix_caching: + # Memory-only mode: prefix caching on a hybrid model without a Mamba + # cache budget deduplicates identical KV prefixes for memory savings, + # but does NOT cache Mamba recurrent state. Prefill skipping is + # therefore disabled (prefix_skip_tokens is forced to 0) and every + # token is recomputed, so results stay correct -- but the main latency + # benefit of prefix caching is forgone. Warn so a user who expected + # full caching knows to set prefix_caching_mamba_gb. + logging.warning( + "enable_prefix_caching is set on a hybrid (Mamba) model but " + "prefix_caching_mamba_gb is not configured (got %r). Running in " + "memory-only mode: identical KV prefixes are deduplicated for " + "memory savings, but Mamba state caching and prefill skipping are " + "disabled (every token is recomputed). Set prefix_caching_mamba_gb " + "> 0 to enable full prefix caching.", + self.config.prefix_caching_mamba_gb, + ) + + # MTP speculative decoding: persistent buffer for decoder hidden states. + # Only needed for block-scope CUDA graphs, where the Python assignment in + # forward() runs only during graph capture. Using copy_() into a fixed + # buffer ensures every batch-size graph replay writes to the same GPU + # address. Sized to max_tokens; only [:actual_tokens] is valid each step. + if ( + self.num_speculative_tokens > 0 + and self.inference_cuda_graph_scope == InferenceCudaGraphScope.block + ): + self.mtp_decoder_hidden_states = torch.empty( + self.max_tokens, + 1, + self.hidden_size, + device=torch.cuda.current_device(), + dtype=self.params_dtype, + ) # Reset tensor-related metadata. self.reset_metadata() @@ -1595,15 +1707,32 @@ def _allocate_mamba_cache(self, mamba_gb: float) -> None: ssm_size = _math.prod(self.mamba_ssm_states_shape) * self.mamba_ssm_states_dtype.itemsize per_slot_bytes = self.num_mamba_layers * (conv_size + ssm_size) total_bytes = int(mamba_gb * 1024**3) - max_slots = total_bytes // per_slot_bytes + + # MambaSlotAllocator allocates two GPU buffer families with the same + # per-slot footprint, both of which must fit in this budget: + # - "durable" cache: self.ssm_states / self.conv_states, sized to + # `max_slots` slots (computed below). + # - "scratch" buffers: self.intermediate_ssm_out / self.intermediate_conv_out, + # fixed CUDA-graph-safe staging for intermediate-state + # extraction, sized to the per-step token-budget cap + # `scratch_slots` = max_mamba_intermediate_states_per_step. + # The scratch is not part of the durable cache but consumes the same + # per-slot bytes, so reserve it from the budget up front before sizing the + # durable cache; otherwise total usage silently exceeds mamba_gb (and can + # OOM) when scratch_slots > max_slots. + scratch_slots = self.max_mamba_intermediate_states_per_step + scratch_bytes = scratch_slots * per_slot_bytes + max_slots = (total_bytes - scratch_bytes) // per_slot_bytes # durable slots if max_slots < 1: - logging.warning( - "Mamba cache budget (%.3f GB) too small for even 1 slot " - "(need %.3f GB per slot). Mamba caching disabled.", - mamba_gb, - per_slot_bytes / 1024**3, + raise ValueError( + f"Mamba prefix cache budget (prefix_caching_mamba_gb={mamba_gb:.4g} GB) " + f"is too small. The CUDA-graph extraction scratch reserves " + f"{scratch_bytes / 1024**3:.4g} GB ({scratch_slots} slots x " + f"{per_slot_bytes / 1024:.1f} KB/slot), leaving room for " + f"fewer than one durable cache slot. Increase prefix_caching_mamba_gb " + f"to at least {(scratch_bytes + per_slot_bytes) / 1024**3:.4g} GB, or " + f"reduce max_tokens." ) - return self.mamba_slot_allocator = MambaSlotAllocator( context=self, @@ -1619,9 +1748,14 @@ def _allocate_mamba_cache(self, mamba_gb: float) -> None: ) logging.info( - "Mamba prefix cache: %d slots (%.3f GB), per-slot %.1f KB", + "Mamba prefix cache: %d durable slots (%.3f GB) + %d scratch slots " + "(%.3f GB) = %.3f GB total within %.3f GB budget, per-slot %.1f KB", max_slots, max_slots * per_slot_bytes / 1024**3, + scratch_slots, + scratch_bytes / 1024**3, + (max_slots + scratch_slots) * per_slot_bytes / 1024**3, + mamba_gb, per_slot_bytes / 1024, ) @@ -2049,7 +2183,9 @@ def initialize_attention_state( *, construct_graph_dimensions: Optional[InferenceBatchDimensions] = None, is_expert_parallel_dummy_cuda_graph_step: bool = False, - ) -> None: + transfer_bookkeeping_to_gpu: bool = True, + record_bookkeeping_done_event: bool = False, + ) -> Optional[torch.cuda.Event]: """Initialize attention state so that every layer can use it. Args: @@ -2057,8 +2193,17 @@ def initialize_attention_state( The graph config to use for constructing the cuda graphs. is_expert_parallel_dummy_cuda_graph_step (bool): Whether this is a dummy expert model parallel step. - Return: - None. + transfer_bookkeeping_to_gpu (bool): Whether to publish the prepared + CPU bookkeeping snapshot to GPU before returning. Legacy + callers publish immediately; async scheduling binds the GPU + views here and publishes their values later. + record_bookkeeping_done_event (bool): Whether to record an event + after the bookkeeping H2D transfer. + + Returns: + Optional[torch.cuda.Event]: Event marking bookkeeping H2D + completion, or `None` when no event was requested or no + transfer was performed. """ # Launch deferred Mamba GPU ops first (state zeroing/restore) so they # overlap with the CPU work below. These are non-blocking GPU kernels. @@ -2296,12 +2441,17 @@ def initialize_attention_state( # No-op when the queue is already empty (regular non-warmup steps). self._execute_pending_mamba_ops() - # Run the H2D transfer here so callers that bypass the controller - # (e.g. unit tests that call `model.forward()` directly after - # `initialize_attention_state()`) see populated GPU bookkeeping. The - # text-generation controller still calls `transfer_bookkeeping_to_gpu` - # explicitly; that second call is a cheap idempotent re-copy. - self.transfer_bookkeeping_to_gpu() + # Record whether this step produces real output — false on CUDA-graph + # capture (warmup) or dummy EP steps. Used by transfer_bookkeeping_to_gpu + # to publish real_token_count=0 so MoE routing masks all padding tokens. + self._bookkeeping_no_real_work = ( + construct_graph_dimensions is not None or is_expert_parallel_dummy_cuda_graph_step + ) + + # Preserve the existing behavior for callers that do not publish explicitly. + if transfer_bookkeeping_to_gpu: + return self.transfer_bookkeeping_to_gpu(record_done_event=record_bookkeeping_done_event) + return None def _execute_pending_mamba_ops(self) -> None: """Execute Mamba GPU operations deferred from add_request() / update_requests(). @@ -2327,18 +2477,33 @@ def _execute_pending_mamba_ops(self) -> None: self.mamba_ssm_states[:, indices] = 0.0 self._pending_mamba_zeros.clear() - def transfer_bookkeeping_to_gpu(self) -> None: + def transfer_bookkeeping_to_gpu( + self, skip_token_input_ids: bool = False, record_done_event: bool = False + ) -> Optional[torch.cuda.Event]: """Batch transfer CPU bookkeeping state to GPU staging buffers. - Called after initialize_attention_state() and before the forward pass. - All copies use non_blocking=True with pinned CPU memory. CUDA stream - ordering guarantees the forward pass sees completed transfers. + Legacy steps call this from initialize_attention_state(). Async + scheduling instead delays publication until after preparation and the + GPU sample-to-input copy. Legacy transfers block because the pinned CPU + source is re-staged in place. Async scheduling requests an event-tracked + non-blocking copy and synchronizes that event before reusing the source. The bookkeeping fields are backed by one contiguous pinned CPU buffer - and one contiguous GPU buffer; a single cudaMemcpyAsync suffices. - Request-level staging slots are refreshed from the persistent CPU - tensors immediately before the H2D (GPU reads them at `[:n_active]` + and one contiguous GPU buffer; a single memcpy covers the whole + transfer. Request-level staging slots are refreshed from the persistent + CPU tensors immediately before the H2D (GPU reads them at `[:n_active]` while CPU bookkeeping keeps them at `[paused_count:total_count)`). + + Args: + skip_token_input_ids (bool): If true, leave + `gpu_view.token_to_input_ids` unchanged while copying the rest + of the bookkeeping buffer. + record_done_event (bool): Whether to record and return an event after + an asynchronous bookkeeping transfer. + + Returns: + Optional[torch.cuda.Event]: Event marking H2D completion, or `None` + when no event was requested. """ n_active = self.total_request_count - self.paused_request_count active_slice = slice(self.paused_request_count, self.total_request_count) @@ -2372,11 +2537,32 @@ def transfer_bookkeeping_to_gpu(self) -> None: self._staging_request_query_lengths[n_active:padded_active] = 0 self._staging_request_kv_length_offsets[n_active:padded_active] = 0 - # Coalesced H2D: one cudaMemcpyAsync for the entire bookkeeping buffer. + # Real (unpadded) token count for this step. CUDA-graph replay pads + # the token dim to a captured size; MoE routing reads this on GPU and + # rewrites padding rows' routing entries to -1 so they don't go to + # any expert. Set to 0 on CUDA-graph capture / dummy EP steps so + # every row gets masked out. + self._staging_real_token_count[0] = ( + 0 if self._bookkeeping_no_real_work else self.batch_dimensions.token_count + ) + + # Coalesced H2D: one copy for the entire bookkeeping buffer. # Copying the whole (max_tokens + max_requests)-sized buffer including # unused slots is cheap (~71 KB total, ~3-5 us on PCIe Gen4) and saves - # 8 redundant launch overheads vs. the prior per-field copies. - self.gpu_view._buf.copy_(self._cpu_bookkeeping_buf, non_blocking=True) + # redundant launch overheads vs. per-field copies. Async scheduling + # decode steps skip token_to_input_ids here because sampled tokens are + # already GPU-resident and copied directly into the GPU input buffer. + if skip_token_input_ids: + token_to_input_ids_offset = ( + self.token_to_input_ids.numel() * self.token_to_input_ids.element_size() + ) + else: + token_to_input_ids_offset = 0 + # Only event-tracked callers may leave the copy in flight; legacy callers + # block before the pinned CPU source can be re-staged. + self.gpu_view._buf[token_to_input_ids_offset:].copy_( + self._cpu_bookkeeping_buf[token_to_input_ids_offset:], non_blocking=record_done_event + ) # MHA metadata GPU views were already bound to state_data in # initialize_attention_state(); the H2D above populates the underlying @@ -2387,6 +2573,54 @@ def transfer_bookkeeping_to_gpu(self) -> None: self.mamba_metadata.load_from_cpu(self._pending_mamba_transfer) self._pending_mamba_transfer = None + done_event = None + if record_done_event: + done_event = self._bookkeeping_h2d_done_event + done_event.record(torch.cuda.current_stream()) + + return done_event + + def copy_async_sched_sample_to_forward( + self, sampled_tokens_cuda: Tensor, sampled_mtp_tokens_cuda: Optional[Tensor] = None + ) -> None: + """Populate GPU input token IDs from sampled CUDA tokens for async scheduling. + + Async scheduling keeps sampled tokens GPU-resident for the next decode + forward. CPU bookkeeping is prepared independently and published later; + this direct GPU copy populates the live input-ID view without waiting + for the sample's CPU copy. + + Args: + sampled_tokens_cuda (Tensor): 1D CUDA tensor containing one sampled + token per active decode request. + sampled_mtp_tokens_cuda (Optional[Tensor]): MTP draft tokens with shape + ``[num_speculative_tokens, active_request_count]``. + """ + active_request_count = self.total_request_count - self.paused_request_count + + if self.num_speculative_tokens > 0: + expected_shape = (self.num_speculative_tokens, active_request_count) + if sampled_mtp_tokens_cuda is None or tuple(sampled_mtp_tokens_cuda.shape) != ( + expected_shape + ): + actual_shape = ( + None if sampled_mtp_tokens_cuda is None else sampled_mtp_tokens_cuda.shape + ) + raise RuntimeError( + f"Expected MTP draft token shape {expected_shape}, got {actual_shape}." + ) + + tokens_per_request = self.num_speculative_tokens + 1 + token_count = active_request_count * tokens_per_request + grouped_tokens = self.gpu_view.token_to_input_ids[:token_count].view( + active_request_count, tokens_per_request + ) + grouped_tokens[:, 0].copy_(sampled_tokens_cuda, non_blocking=True) + if sampled_mtp_tokens_cuda is not None: + grouped_tokens[:, 1:].copy_(sampled_mtp_tokens_cuda.transpose(0, 1), non_blocking=True) + if token_count < self.padded_active_token_count: + self.gpu_view.token_to_input_ids[token_count : self.padded_active_token_count].zero_() + def reset_tensors(self) -> None: """Fill all bookkeeping tensors with sentinel values.""" @@ -2413,19 +2647,34 @@ def reset_tensors(self) -> None: self.token_to_block_idx.fill_(-1) self.token_to_local_position_within_kv_block.fill_(0) - def reset_metadata(self) -> None: + def reset_metadata( + self, preserve_prefix_cache: bool = False, *, preserve_counters: bool = False + ) -> None: """Reset all bookkeeping state: counters, block allocator, attention/mamba state. This must be called after ``initialize_all_tensors()`` and after any suspend/resume cycle to bring the context back to a clean state. + + Args: + preserve_prefix_cache: When True, keep the KV block allocator's prefix-cache + state (hash index, ref counts, cached blocks) intact. Used by the idle + ``dummy_forward`` path, which only needs to clear the transient one-token + step state -- wiping the allocator there would destroy cross-request prefix + reuse for any subsequent request (the engine idles between requests at low + concurrency, especially with EP > 1). + preserve_counters: When True, keep engine-step, prefix-cache clock, + prefill-token, and async-scheduling counters intact. """ + # There is no prefix-cache state to preserve when caching is disabled. + preserve_prefix_cache = preserve_prefix_cache and self.enable_prefix_caching # Reset request/token counts. self.total_request_count = 0 self.active_token_count = 0 - self.lifetime_prefill_token_count = 0 - self.async_sched_step_count = 0 - self.async_sched_compaction_step_count = 0 + if not preserve_counters: + self.lifetime_prefill_token_count = 0 + self.async_sched_step_count = 0 + self.async_sched_compaction_step_count = 0 self.paused_request_count = 0 self.batch_dimensions = InferenceBatchDimensions( token_count=0, prefill_req_count=0, decode_req_count=0 @@ -2441,7 +2690,8 @@ def reset_metadata(self) -> None: # Reset attention, mamba, and block allocator state. self.reset_attention_state() self.reset_mamba_state() - self.kv_block_allocator.reset() + if not preserve_prefix_cache: + self.kv_block_allocator.reset() self.request_to_kv_block_ids.fill_(-1) # Reset chunked prefill state @@ -2453,7 +2703,9 @@ def reset_metadata(self) -> None: token_count=0, prefill_req_count=0, decode_req_count=0 ) - def reset(self) -> None: + def reset( + self, preserve_prefix_cache: bool = False, *, preserve_counters: bool = False + ) -> None: """Reset entire context. This method does: @@ -2464,17 +2716,28 @@ def reset(self) -> None: This method is useful after cuda graph warmup iterations, where the context's memory buffer is referenced by the cuda graph system and cannot be deallocated. + + Args: + preserve_prefix_cache: When True, keep the KV and Mamba prefix-cache + state (hash indices and cached blocks/slots) intact. Used by + the idle ``dummy_forward`` path so an idle step between requests does + not destroy cross-request prefix reuse. + preserve_counters: When True, keep engine-step, prefix-cache clock, + prefill-token, and async-scheduling counters intact. """ + # There is no prefix-cache state to preserve when caching is disabled. + preserve_prefix_cache = preserve_prefix_cache and self.enable_prefix_caching self.reset_tensors() - self.reset_metadata() + self.reset_metadata( + preserve_prefix_cache=preserve_prefix_cache, preserve_counters=preserve_counters + ) - # Reset lifetime counters (not reset in reset_metadata, which is also - # called during suspend/resume where these must persist). - self.step_count = 0 - self.prefix_cache_lru_clock = 0 + if not preserve_counters: + self.step_count = 0 + self.prefix_cache_lru_clock = 0 - # Reset Mamba cache state - if self.mamba_slot_allocator is not None: + # Reset Mamba cache state. + if not preserve_prefix_cache and self.mamba_slot_allocator is not None: self.mamba_slot_allocator.reset() def current_input_and_position_ids( @@ -2564,13 +2827,46 @@ def last_token_logits(self, logits: Tensor) -> Tensor: ) return logits.squeeze(0)[self.active_logit_idxs[: self.num_last_token_logits], :] + def _find_mamba_match_count( + self, req: DynamicInferenceRequest, start_block: int, end_block: int + ) -> int: + """Find the farthest cached Mamba state within a chunk-local block range. + + Mamba state restore is only valid for blocks that the current chunk also + assigns from the KV cache. Chunked prefill can schedule a prompt prefix + that is shorter than the farthest cached full-prompt Mamba boundary, so + this helper intentionally uses the same block domain as KV matching. + """ + if self.mamba_slot_allocator is None or not req.precomputed_block_hashes: + return 0 + + end_block = min(end_block, len(req.precomputed_block_hashes)) + if start_block >= end_block: + return 0 + + mamba_map = self.mamba_slot_allocator.hash_to_block_id + hashes = req.precomputed_block_hashes[start_block:end_block] + for i in range(len(hashes) - 1, -1, -1): + if hashes[i] in mamba_map: + return i + 1 + return 0 + def _compute_prefix_match( - self, req: DynamicInferenceRequest, prefill_chunk_length: int + self, + req: DynamicInferenceRequest, + prefill_chunk_length: int, + record_mamba_match: bool = False, ) -> Tuple[list, int, int, int, int, int]: """Compute prefix match results and skip counts for a request chunk. Shared by check_availability (budget checks) and add_request (execution). + Args: + req: Request being scheduled. + prefill_chunk_length: Number of prompt tokens considered in this chunk. + record_mamba_match: If True, store the chunk-local executable Mamba + match count on the request for diagnostics/tests. + Returns: Tuple of (matched_block_ids, num_blocks_from_pool, already_allocated_blocks, overall_required_blocks, @@ -2609,7 +2905,11 @@ def _compute_prefix_match( # Only applies to the first chunk (finished == 0); continuation chunks # already had Mamba state restored during the first chunk. if self.is_hybrid_model and self.mamba_slot_allocator is not None and finished == 0: - num_mamba_matched = getattr(req, '_mamba_num_matched_blocks', 0) + num_mamba_matched = self._find_mamba_match_count( + req, already_allocated_blocks, already_allocated_blocks + num_matched + ) + if record_mamba_match: + req._mamba_num_matched_blocks = num_mamba_matched assert ( num_mamba_matched <= num_matched ), f"Mamba match ({num_mamba_matched}) > KV match ({num_matched})" @@ -2629,6 +2929,8 @@ def _compute_prefix_match( else: prefix_skip_tokens = 0 elif self.is_hybrid_model and finished == 0: + if record_mamba_match: + req._mamba_num_matched_blocks = 0 prefix_skip_tokens = 0 # Clamp so that effective_prefill_chunk_length >= 2 when possible. @@ -2663,14 +2965,27 @@ def check_availability(self, req: DynamicInferenceRequest) -> Tuple[bool, bool, self.total_request_count < self.max_requests and self.paused_request_count == 0 ) - (_, num_blocks_from_pool, _, _, _, effective_prefill_chunk_length) = ( + (matched_block_ids, num_blocks_from_pool, _, _, _, effective_prefill_chunk_length) = ( self._compute_prefix_match(req, req.remaining_prompt_length) ) request_tokens_can_be_added = ( self.active_token_count + effective_prefill_chunk_length <= self.max_tokens ) - kv_cache_available = self.kv_block_allocator.is_memory_available(num_blocks_from_pool) + # add_request pins the matched blocks before allocating. Only matches that + # are currently evictable (ref_count == 0) count against the evictable + # pool; matches already pinned by another in-flight request are not in + # get_evictable_block_count() and pinning them frees nothing. Reserve only + # the ref_count == 0 matches so availability is not under-reported. + potential_matched_count = 0 + if matched_block_ids: + matched_tensor = torch.tensor(matched_block_ids, dtype=torch.int32, device='cpu') + potential_matched_count = int( + (self.kv_block_allocator.block_ref_counts[matched_tensor] == 0).sum() + ) + kv_cache_available = self.kv_block_allocator.is_memory_available( + num_blocks_from_pool, potential_matched_count=potential_matched_count + ) return request_can_be_added, request_tokens_can_be_added, kv_cache_available def _find_kv_match_count( @@ -2755,31 +3070,47 @@ def add_request( overall_required_blocks, prefix_skip_tokens, effective_prefill_chunk_length, - ) = self._compute_prefix_match(req, prefill_chunk_length) + ) = self._compute_prefix_match(req, prefill_chunk_length, record_mamba_match=True) num_matched_blocks = len(matched_block_ids) effective_kv_offset = req.finished_chunk_token_count + prefix_skip_tokens - # Track prefix cache hits. + # Track prefix cache hits. num_cached_tokens accumulates across prefill + # chunks: each chunk matches a disjoint block range (start advances with + # finished_chunk_token_count), so a long cached prefix is discovered + # incrementally and must be summed, not overwritten. if num_matched_blocks > 0: self.prefix_cache_hits += 1 self.prefix_cache_blocks_matched += num_matched_blocks + req.num_cached_tokens += num_matched_blocks * self.block_size_tokens # Slice tokens to skip matched prefix this_round_tokens = req.remaining_prompt_tokens[prefix_skip_tokens:prefill_chunk_length] - new_block_ids = None - if num_blocks_from_pool > 0: - new_block_ids = self.kv_block_allocator.allocate_memory_blocks(num_blocks_from_pool) - if new_block_ids is None or len(new_block_ids) != num_blocks_from_pool: - raise BlockOverflowError(req.request_id) - - # Increment ref counts and update timestamps for matched (shared) blocks + # Pin matched (shared) blocks BEFORE allocation. allocate_memory_blocks() + # may trigger LRU eviction, and a matched block still at ref_count == 0 + # would be an eviction candidate — descendant-first LRU could evict a + # matched leaf and immediately reuse its ID for a new block, leaving the + # block table with a duplicate ID and a dangling parent. Incrementing ref + # counts first removes matched blocks from the evictable set (see + # evict_lru_blocks / get_evictable_block_count), so eviction falls back to + # genuinely unused cached blocks. + matched_tensor = None if num_matched_blocks > 0: matched_tensor = torch.tensor(matched_block_ids, dtype=torch.int32, device='cpu') self.kv_block_allocator.block_ref_counts[matched_tensor] += 1 if self.prefix_caching_eviction_policy == PrefixCachingEvictionPolicy.LRU: self.kv_block_allocator.update_timestamps(matched_tensor) + new_block_ids = None + if num_blocks_from_pool > 0: + new_block_ids = self.kv_block_allocator.allocate_memory_blocks(num_blocks_from_pool) + if new_block_ids is None or len(new_block_ids) != num_blocks_from_pool: + # Roll back the pin so a failed add does not leak ref counts on + # the matched blocks (which would make them permanently unevictable). + if matched_tensor is not None: + self.kv_block_allocator.block_ref_counts[matched_tensor] -= 1 + raise BlockOverflowError(req.request_id) + # Note that we decremented the total_request_count for the chunked prefill request # in update_requests, so setting current_id to the total_request_count will again # make the last request the continuing chunked prefill request if one exists. @@ -2881,8 +3212,14 @@ def _register_range(start: int, end: int): return block_ids_to_hash = self.request_to_kv_block_ids[current_id][start:end].tolist() block_hashes_slice = req.precomputed_block_hashes[start:end] + # Parent hash of block k is the hash of block k-1 in the chain; + # block 0 is a root (parent hash 0). Enables LRU eviction to keep + # parents cached until their children are gone. + parent_hashes_slice = [ + req.precomputed_block_hashes[k - 1] if k > 0 else 0 for k in range(start, end) + ] self.kv_block_allocator.register_kv_block_hashes( - block_ids_to_hash, block_hashes_slice + block_ids_to_hash, block_hashes_slice, parent_hashes_slice ) # Range 1: prior-chunk partial block that this chunk just completed @@ -2907,23 +3244,28 @@ def _register_range(start: int, end: int): else: self._pending_mamba_zeros.append(mamba_idx) - # compute_and_store_offsets sets both CPU state (hash_to_block_id, - # _eos_cache_block_id_gpu) and GPU staging buffers. Runs immediately - # because commit_intermediate_states() reads the CPU state after the - # forward pass. - if self.mamba_slot_allocator is not None: - self.mamba_slot_allocator.compute_and_store_offsets( - req, - current_id, - prefix_skip_tokens, - prefill_chunk_length, - num_matched_blocks, - matched_block_ids, - overall_required_blocks, - ) + # compute_and_store_offsets sets CPU state + GPU staging buffers that + # commit_intermediate_states() consumes after the forward pass. Run it for + # EVERY prefill chunk (not just the first): the last complete block of a + # multi-chunk prompt falls in a continuation chunk, and caching its Mamba + # state is precisely what lets a later turn skip prefill on a hybrid model. + # Mamba slot allocation / state restore above stays first-chunk-only. + if self.is_hybrid_model and self.mamba_slot_allocator is not None: + self.mamba_slot_allocator.compute_and_store_offsets( + req, + current_id, + prefix_skip_tokens, + prefill_chunk_length, + num_matched_blocks, + matched_block_ids, + overall_required_blocks, + ) self.active_token_count += effective_prefill_chunk_length self.lifetime_prefill_token_count += effective_prefill_chunk_length + if self.enable_prefix_caching: + self.prefix_cache_prefill_computed_tokens += effective_prefill_chunk_length + self.prefix_cache_prefill_skipped_tokens += prefix_skip_tokens self.total_request_count += 1 self.num_prefill_requests += 1 @@ -3239,41 +3581,59 @@ def evict_overflow_paused_requests( return evict_request_ids - def prepare_requests(self, new_tokens: Tensor) -> None: - """Speculatively prepare active decode requests for the next forward pass. + def _get_async_sched_rows_requiring_new_block(self) -> Tensor: + """Return active request rows that need a block during the next prepare. - Async scheduling only supports decode-only steps with no pause, - evict, or resume lifecycle changes. If preparing the next token would - require one of those lifecycle changes, this method raises and the caller - should treat async scheduling as unsupported for that workload. + Returns: + Tensor: Boolean mask over active request rows. + """ + active_slice = slice(self.paused_request_count, self.total_request_count) + tokens_per_request = self.num_speculative_tokens + 1 + return ( + self.request_last_kv_block_offset[active_slice] + tokens_per_request + >= self.block_size_tokens + ) - Args: - new_tokens (Tensor): Newly sampled token for each active request. + def can_prepare_requests(self) -> bool: + """Return whether requests can be prepared without lifecycle changes. + + Returns: + bool: Whether all requests are active decode requests and the active + KV-block pool can satisfy the exact next-step allocation demand. """ - if new_tokens.is_cuda: - new_tokens = new_tokens.cpu() + if self.num_prefill_requests != 0 or self.paused_request_count != 0: + return False + + rows_requiring_new_block = self._get_async_sched_rows_requiring_new_block() + num_new_blocks = int(rows_requiring_new_block.sum().item()) + return num_new_blocks <= self.kv_block_allocator.get_active_avail() + + def prepare_requests(self) -> None: + """Speculatively prepare active decode requests for the next forward pass. + Async scheduling only supports decode-only steps with no pause, + evict, or resume lifecycle changes. If preparation cannot allocate the + required KV blocks without a lifecycle change, this method raises. The + prepared decode layout establishes the active token count. + """ active_request_count = self.total_request_count - self.paused_request_count - if self.num_speculative_tokens != 0: - raise RuntimeError("Async scheduling does not support speculative tokens.") if self.num_prefill_requests != 0: raise RuntimeError("Async scheduling only supports decode-only steps.") if self.paused_request_count != 0: raise RuntimeError("Async scheduling does not support paused requests.") - if new_tokens.numel() != active_request_count: - raise RuntimeError( - f"Expected {active_request_count} new tokens, got {new_tokens.numel()}." - ) if active_request_count == 0: self.active_token_count = 0 return active_slice = slice(0, active_request_count) - rows_requiring_new_block = ( - self.request_last_kv_block_offset[active_slice] >= self.block_size_tokens - 1 - ) - num_new_blocks = rows_requiring_new_block.sum().item() + tokens_per_request = self.num_speculative_tokens + 1 + last_block_offsets = self.request_last_kv_block_offset[active_slice] + token_offsets = self._async_sched_token_offsets + rows_requiring_new_block = self._get_async_sched_rows_requiring_new_block() + num_new_blocks = int(rows_requiring_new_block.sum().item()) + + block_ids = None if num_new_blocks > 0: active_block_count_avail = self.kv_block_allocator.get_active_avail() if num_new_blocks > active_block_count_avail: @@ -3283,53 +3643,123 @@ def prepare_requests(self, new_tokens: Tensor) -> None: if block_ids is None: raise RuntimeError("Async scheduling cannot evict requests to allocate new blocks.") + self.active_token_count = active_request_count * tokens_per_request + active_token_slice = slice(0, self.active_token_count) + grouped_token_block_ids = self.token_to_block_idx[active_token_slice].view( + active_request_count, tokens_per_request + ) + grouped_token_block_ids.copy_(self.request_last_kv_block_id[active_slice, None]) + + if block_ids is not None: row_idx = torch.nonzero(rows_requiring_new_block, as_tuple=True)[0] col_idx = self.request_kv_block_counts[row_idx] self.request_to_kv_block_ids[row_idx, col_idx] = block_ids self.request_kv_block_counts[row_idx] += 1 self.request_last_kv_block_id[row_idx] = block_ids + grouped_token_block_ids[row_idx] = torch.where( + last_block_offsets[row_idx, None] + 1 + token_offsets[None, :] + >= self.block_size_tokens, + block_ids[:, None], + grouped_token_block_ids[row_idx], + ) self.request_kv_length_offsets[active_slice].add_(self.request_query_lengths[active_slice]) - self.request_query_lengths[active_slice].fill_(1) - + self.request_query_lengths[active_slice].fill_(tokens_per_request) self.request_last_kv_block_offset[active_slice] = ( - self.request_last_kv_block_offset[active_slice] + 1 + last_block_offsets + tokens_per_request ) % self.block_size_tokens - self.active_token_count = active_request_count - self.token_to_input_ids[:active_request_count] = new_tokens - self.token_to_pos_ids[:active_request_count] = self.request_kv_length_offsets[active_slice] - self.token_to_request_idx[:active_request_count] = torch.arange( - active_request_count, device='cpu' + token_positions = ( + self.request_kv_length_offsets[active_slice, None] + token_offsets[None, :] ) - self.token_to_position_in_request[:active_request_count] = self.token_to_pos_ids[ - :active_request_count + token_request_idxs = torch.arange(active_request_count, device='cpu').repeat_interleave( + tokens_per_request + ) + self.token_to_pos_ids[active_token_slice] = token_positions.flatten() + self.token_to_request_idx[active_token_slice] = token_request_idxs + self.token_to_position_in_request[active_token_slice] = self.token_to_pos_ids[ + active_token_slice ] - self.token_to_local_position_within_kv_block[:active_request_count] = ( - self.token_to_pos_ids[:active_request_count] % self.block_size_tokens + self.token_to_local_position_within_kv_block[active_token_slice] = ( + self.token_to_pos_ids[active_token_slice] % self.block_size_tokens ) - self.token_to_block_idx[:active_request_count] = self.request_last_kv_block_id[active_slice] - def resolve_requests(self, active_requests_mask: Tensor) -> Tensor: + def commit_sampled_tokens( + self, sampled_tokens_cpu: Tensor, sampled_mtp_tokens_cpu: Optional[Tensor] = None + ) -> None: + """Commit sampled CPU token IDs to the prepared request state. + + This establishes the post-resolution active token count and populates + the CPU input-ID staging rows in survivor order. Overlapped async + scheduling has already copied the same samples into the live GPU input + view for the speculative forward. + + Args: + sampled_tokens_cpu (Tensor): Sampled CPU token for each active request. + sampled_mtp_tokens_cpu (Optional[Tensor]): MTP draft tokens with shape + ``[num_speculative_tokens, active_request_count]``. + """ + assert sampled_tokens_cpu.device == torch.device( + 'cpu' + ), "Sampled tokens must be on the CPU before they are committed." + + if sampled_mtp_tokens_cpu is not None: + assert sampled_mtp_tokens_cpu.device == torch.device( + 'cpu' + ), "MTP draft tokens must be on the CPU before they are committed." + + active_request_count = self.total_request_count - self.paused_request_count + if sampled_tokens_cpu.numel() != active_request_count: + raise RuntimeError( + f"Expected {active_request_count} new tokens, got {sampled_tokens_cpu.numel()}." + ) + + expected_mtp_shape = (self.num_speculative_tokens, active_request_count) + if self.num_speculative_tokens == 0: + if sampled_mtp_tokens_cpu is not None and sampled_mtp_tokens_cpu.numel() != 0: + raise RuntimeError( + "Received MTP draft tokens when speculative decoding is disabled." + ) + else: + if sampled_mtp_tokens_cpu is None or tuple(sampled_mtp_tokens_cpu.shape) != ( + expected_mtp_shape + ): + actual_shape = ( + None if sampled_mtp_tokens_cpu is None else sampled_mtp_tokens_cpu.shape + ) + raise RuntimeError( + f"Expected MTP draft token shape {expected_mtp_shape}, got {actual_shape}." + ) + + tokens_per_request = self.num_speculative_tokens + 1 + active_token_count = active_request_count * tokens_per_request + self.active_token_count = active_token_count + grouped_tokens = self.token_to_input_ids[:active_token_count].view( + active_request_count, tokens_per_request + ) + grouped_tokens[:, 0] = sampled_tokens_cpu + if sampled_mtp_tokens_cpu is not None: + grouped_tokens[:, 1:] = sampled_mtp_tokens_cpu.transpose(0, 1) + + def resolve_requests(self, active_requests_mask: Tensor) -> Tuple[Tensor, Tensor]: """Resolve finished requests after an async scheduling forward pass. - Async scheduling supports only request completion. The active request rows - and current decode-token rows are compacted in survivor order so any - following legacy or async scheduling step sees a consistent context. + Prefill requests transition to decode during resolution. Request rows use + the same hole-filling order as ``update_requests`` so seeded sampling stays + consistent with legacy scheduling. Token tensors and the active token + count are left untouched; prepare rebuilds derived token metadata and the + controller commits sampled input IDs after resolution. Args: active_requests_mask (Tensor): 1D mask marking requests that remain active. Returns: - Tensor: Request IDs for requests that finished during resolution. + Tuple[Tensor, Tensor]: Request IDs that finished and source row indices + for surviving requests in their resolved destination order. """ if active_requests_mask.is_cuda: active_requests_mask = active_requests_mask.cpu() - if self.num_speculative_tokens != 0: - raise RuntimeError("Async scheduling does not support speculative tokens.") - if self.num_prefill_requests != 0: - raise RuntimeError("Async scheduling only supports decode-only steps.") if self.paused_request_count != 0: raise RuntimeError("Async scheduling does not support paused requests.") @@ -3340,29 +3770,38 @@ def resolve_requests(self, active_requests_mask: Tensor) -> Tensor: f"got {active_requests_mask.numel()}." ) - survivor_idxs = torch.nonzero(active_requests_mask == 1, as_tuple=True)[0] + self.num_prefill_requests = 0 + self.request_in_prefill_status_tensor[self.request_in_prefill_status_tensor == 1] = 0 + finished_idxs = torch.nonzero(active_requests_mask == 0, as_tuple=True)[0] finished_request_ids = self.request_ids[finished_idxs].clone() + active_request_count = int(active_requests_mask.sum().item()) + survivor_idxs = torch.arange(active_request_count, device='cpu') + finished_idxs_on_left = torch.nonzero( + active_requests_mask[:active_request_count] == 0, as_tuple=True + )[0] + active_idxs_on_right = ( + torch.nonzero(active_requests_mask[active_request_count:] == 1, as_tuple=True)[0] + + active_request_count + ) + assert finished_idxs_on_left.numel() == active_idxs_on_right.numel() + survivor_idxs[finished_idxs_on_left] = active_idxs_on_right + self.reset_attention_state() if finished_idxs.numel() > 0: self.release_memory_blocks_from_request_indexes(finished_idxs) - active_request_count = survivor_idxs.numel() if active_request_count == 0: self.request_to_kv_block_ids.fill_(-1) self.total_request_count = 0 - self.active_token_count = 0 self.reset_mamba_state() - return finished_request_ids + return finished_request_ids, survivor_idxs dst_idxs = torch.arange(active_request_count, device='cpu') if not torch.equal(survivor_idxs, dst_idxs): self.request_kv_length_offsets[dst_idxs] = self.request_kv_length_offsets[survivor_idxs] - self.request_in_prefill_status_tensor[dst_idxs] = self.request_in_prefill_status_tensor[ - survivor_idxs - ] self.request_query_lengths[dst_idxs] = self.request_query_lengths[survivor_idxs] self.request_output_lengths[dst_idxs] = self.request_output_lengths[survivor_idxs] self.request_ids[dst_idxs] = self.request_ids[survivor_idxs] @@ -3372,25 +3811,18 @@ def resolve_requests(self, active_requests_mask: Tensor) -> Tensor: self.request_last_kv_block_offset[dst_idxs] = self.request_last_kv_block_offset[ survivor_idxs ] + if self.is_hybrid_model: + self.mamba_metadata.request_to_mamba_state_idx[dst_idxs] = ( + self.mamba_metadata.request_to_mamba_state_idx[survivor_idxs] + ) for metadata_tensor in self.request_metadata.values(): metadata_tensor[dst_idxs] = metadata_tensor[survivor_idxs] - - self.token_to_input_ids[dst_idxs] = self.token_to_input_ids[survivor_idxs] - self.token_to_pos_ids[dst_idxs] = self.token_to_pos_ids[survivor_idxs] - self.token_to_block_idx[dst_idxs] = self.token_to_block_idx[survivor_idxs] - self.token_to_local_position_within_kv_block[dst_idxs] = ( - self.token_to_local_position_within_kv_block[survivor_idxs] - ) - self.token_to_position_in_request[dst_idxs] = self.token_to_position_in_request[ - survivor_idxs - ] - - self.token_to_request_idx[:active_request_count] = dst_idxs stale_slice = slice(active_request_count, old_active_request_count) self.request_to_kv_block_ids[stale_slice] = -1 + if self.is_hybrid_model: + self.mamba_metadata.request_to_mamba_state_idx[stale_slice] = -1 self.total_request_count = active_request_count - self.active_token_count = active_request_count - return finished_request_ids + return finished_request_ids, survivor_idxs def update_requests( self, @@ -3855,25 +4287,83 @@ def _processed_log_probs( n_active: int, active_query_lengths: Optional[Tensor], sampling: Optional[Sampling], + row_to_request: Optional[Tensor] = None, ) -> Tensor: - """Sample the logprobs if desired.""" + """Calculate raw or sampling-processed per-row log probabilities. + + Args: + logits (Tensor): Raw logits with shape `[num_rows, vocab_size]`. + n_active (int): Number of active requests represented by the rows. + active_query_lengths (Optional[Tensor]): CPU token counts used to + map prefill rows to active requests, or `None` for decode. + sampling (Optional[Sampling]): Backend providing processed logprobs. + row_to_request (Optional[Tensor]): Explicit CPU mapping from each + logit row to an active request. + + Returns: + Tensor: Per-row log probabilities over the vocabulary. + """ if self.config.logprobs_mode == "raw_logprobs": return F.log_softmax(logits, dim=-1) assert sampling is not None, "processed_logprobs requires a sampling backend" # Map each logits row to its active request. - request_idx = torch.arange(n_active, device=logits.device) - row_to_request = ( - request_idx - if active_query_lengths is None - else request_idx.repeat_interleave(active_query_lengths) - ) - md = self.active_request_metadata - temperature = md["temperature"][:n_active].to(logits.device, torch.float32)[row_to_request] - top_k = md["top_k"][:n_active].to(logits.device, torch.long)[row_to_request] - top_p = md["top_p"][:n_active].to(logits.device, torch.float32)[row_to_request] - return sampling.log_probs_kernel(logits, temperature, top_k, top_p) + if row_to_request is None and active_query_lengths is not None: + row_to_request = torch.arange(n_active).repeat_interleave(active_query_lengths) + return sampling.log_probs_kernel(logits, self, token_to_request_index=row_to_request) + + def calculate_log_probs_tensors( + self, + logits: Tensor, + new_tokens: Tensor, + only_last_token_logits: Optional[bool] = False, + sampling: Optional[Sampling] = None, + row_to_request: Optional[Tensor] = None, + ) -> Tuple[Tensor, Tensor]: + """Calculate selected-token and full-distribution log probabilities. + + Args: + logits (Tensor): Raw model output logits with shape + `[1, sequence_length, vocab_size]`. + new_tokens (Tensor): Newly sampled tokens for active requests. + only_last_token_logits (Optional[bool]): Whether logits contain only + each request's final token row. + sampling (Optional[Sampling]): Sampling backend used for processed + log probabilities. + row_to_request (Optional[Tensor]): Explicit CPU mapping from each + logit row to an active request. + + Returns: + Tuple[Tensor, Tensor]: Selected-token log probabilities flattened in + active-token order and the full per-row log-probability tensor. + """ + logits_squeezed = logits.squeeze(0) + n_active = self.total_request_count - self.paused_request_count + + if only_last_token_logits or self.is_decode_only(): + seq_idx = torch.arange(len(new_tokens), dtype=torch.int32, device=logits.device) + active_logits = logits_squeezed[: len(new_tokens)].float() + log_probs = self._processed_log_probs( + active_logits, n_active, None, sampling, row_to_request + ) + return log_probs[seq_idx, new_tokens], log_probs + + logits_squeezed = logits_squeezed.float() + active_slice = slice(self.paused_request_count, self.total_request_count) + active_query_lengths_cpu = self.request_query_lengths[active_slice] + active_query_lengths_gpu = self.gpu_view.request_query_lengths[:n_active] + + # Shift away each request's first prompt token, then insert its sampled token. + active_token_ids = self.gpu_view.token_to_input_ids[: self.active_token_count].roll(-1, 0) + new_token_idx = active_query_lengths_gpu.cumsum(0) - 1 + active_token_ids[new_token_idx] = new_tokens + + log_probs = self._processed_log_probs( + logits_squeezed, n_active, active_query_lengths_cpu, sampling + ) + seq_idx = torch.arange(self.active_token_count, device=log_probs.device) + return log_probs[seq_idx, active_token_ids], log_probs def calculate_log_probs( self, @@ -3898,58 +4388,15 @@ def calculate_log_probs( log_probs (Tensor): Used to compute top n logprobs later if required. """ - # Calculate log_probs (sequence_length x vocab_size) - logits_squeezed = logits.squeeze(0).float() - n_active = self.total_request_count - self.paused_request_count + selected_log_probs, log_probs = self.calculate_log_probs_tensors( + logits, new_tokens, only_last_token_logits=only_last_token_logits, sampling=sampling + ) if only_last_token_logits or self.is_decode_only(): - seq_idx = torch.arange(len(new_tokens), dtype=torch.int32, device=logits.device) - log_probs = self._processed_log_probs( - logits_squeezed[seq_idx], n_active, None, sampling - ) - selected_log_probs = log_probs[seq_idx, new_tokens] return [[lp] for lp in selected_log_probs.tolist()], log_probs - # Get the selected token ids for all tokens. - # We shift the active token window left by one to remove the first prompt token for - # prefill requests and then set the token ids explicitly for the newly generated tokens. - # This is necessary because we calculate the log probs *before* updating the request metadata. - # - # Example (decode & prefill mix): - # - # active_query_lengths: [ 1 | 1 | 2 | 5 ] - # - # new_tokens : [ 52 | 12 | 3 | 86 ] - # - # seq_idx : [ 0 | 1 | 2 3 | 4 5 6 7 8 ] - # - # new_token_idx : [ 0 | 1 | 3 | 8 ] - # - # active_token_ids before left shift: - # : [ 31 | 75 | 45 16 | 90 12 72 24 88 ] - # - # active_token_ids after shift: - # : [ XX | XX | 16 XX | 12 72 24 88 XX ] (XX = undefined) - # - # active_token_ids[new_token_idx] = new_tokens - # : [ 52 | 12 | 16 3 | 12 72 24 88 86 ] - active_token_ids = self.gpu_view.token_to_input_ids[: self.active_token_count].roll(-1, 0) - active_query_lengths = self.gpu_view.request_query_lengths[:n_active] - - new_token_idx = active_query_lengths.cumsum(0) - 1 - active_token_ids[new_token_idx] = new_tokens - - # Compute (possibly processed) log-probs over all active-token rows. - log_probs = self._processed_log_probs( - logits_squeezed, n_active, active_query_lengths, sampling - ) - - # Extract the log probs for only the selected tokens. - # (sequence_length x vocab_size) -> (sequence_length) - seq_idx = torch.arange(self.active_token_count, device=log_probs.device) - selected_log_probs = log_probs[seq_idx, active_token_ids] - - # Split the log probs across request boundaries + active_slice = slice(self.paused_request_count, self.total_request_count) + active_query_lengths = self.request_query_lengths[active_slice] selected_log_probs_list = selected_log_probs.cpu().split( active_query_lengths.tolist(), dim=0 ) diff --git a/megatron/core/inference/contexts/gpu_view.py b/megatron/core/inference/contexts/gpu_view.py index 651d95055da..2066375d19e 100644 --- a/megatron/core/inference/contexts/gpu_view.py +++ b/megatron/core/inference/contexts/gpu_view.py @@ -43,6 +43,9 @@ def __init__( # query_lengths, kv_length_offsets) + 1 int32 (top_k) + 2 float32 # (temperature, top_p) + 1 int32 (active_request_last_token_idxs) = 7 fields. req_4byte_bytes = max_requests * 4 + # Scalar: real (unpadded) token count for the current step. Used by + # MoE routing to mask out CUDA-graph padding tokens. + real_token_count_bytes = 4 # MHA section: 5 fields shared by both graphed and non-graphed MHAMetadata # (only one is active per step, so sharing storage is fine). @@ -74,6 +77,7 @@ def __init__( 3 * tok_int64_bytes + 3 * tok_int32_bytes + 7 * req_4byte_bytes + + real_token_count_bytes + mha_query_lengths_bytes + mha_cu_query_seq_lengths_bytes + mha_kv_seq_lengths_bytes @@ -163,6 +167,12 @@ def __init__( ) off += req_4byte_bytes + # Real (unpadded) token count for the current step. Scalar int32 view. + # MoE routing reads this to skip routing CUDA-graph padding tokens to + # experts. Refreshed each step by transfer_bookkeeping_to_gpu(). + self.real_token_count = self._buf[off : off + real_token_count_bytes].view(torch.int32) + off += real_token_count_bytes + # MHA flash-attention metadata (shared between GraphedMHAMetadata and # NonGraphedMHAMetadata — only one is active per step). self.mha_query_lengths = self._buf[off : off + mha_query_lengths_bytes].view(torch.int32) diff --git a/megatron/core/inference/contexts/kv_block_allocator.py b/megatron/core/inference/contexts/kv_block_allocator.py index d555c925c93..3feb8a0a11d 100644 --- a/megatron/core/inference/contexts/kv_block_allocator.py +++ b/megatron/core/inference/contexts/kv_block_allocator.py @@ -1,5 +1,6 @@ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +import heapq from collections import deque from typing import Callable, Dict, Optional @@ -70,6 +71,24 @@ def __init__( (self.total_count,), dtype=torch.int64, device='cpu' ) + # Persisted prefix-chain bookkeeping for LRU eviction, maintained + # incrementally on register/deregister. Block hashes are + # parent-chained: a cached block that is another cached block's + # parent must not be evicted before its child (see evict_lru_blocks). + # + # block_parent_id[b] = block id of b's parent in the prefix chain, + # or -1 when b is a root block or its parent is not registered. + self.block_parent_id = torch.full( + (self.total_count,), -1, dtype=torch.int64, device='cpu' + ) + # block_child_count[b] = number of currently-registered children of b. + # For a cached block all of its children are cached too, so this + # equals its cached-child count and b is an evictable leaf exactly + # when it reaches 0. + self.block_child_count = torch.zeros( + (self.total_count,), dtype=torch.int64, device='cpu' + ) + # Per-block MoE routing storage (populated when routing replay is enabled) self.block_routing: Dict[int, np.ndarray] = {} @@ -128,13 +147,20 @@ def get_paused_avail(self): """Compute number of paused blocks available.""" return self.paused_count - self.get_paused_used() - def is_memory_available(self, num_blocks: int) -> bool: + def is_memory_available(self, num_blocks: int, potential_matched_count: int = 0) -> bool: """Check if memory blocks are available. Includes both free pool blocks and evictable cached blocks (ref_count == 0). Args: num_blocks (int): Number of blocks to check. + potential_matched_count (int): Number of currently-evictable cached + blocks to subtract from the evictable count because the caller + will pin them before allocating (e.g. prefix-matched blocks that + get their ref counts bumped in add_request). These blocks are + ref_count == 0 now, so they are included in the evictable count, + but they will be protected from eviction, so they cannot supply + the requested ``num_blocks``. Return: (bool) Is memory available? @@ -146,8 +172,8 @@ def is_memory_available(self, num_blocks: int) -> bool: return False if self.prefix_caching_eviction_policy == PrefixCachingEvictionPolicy.REF_ZERO: return False # RZ: no cached blocks to evict - # Also count evictable cached blocks - evictable_count = self.get_evictable_block_count() + # Also count evictable cached blocks, excluding those the caller will pin. + evictable_count = int(self.get_evictable_block_count()) - potential_matched_count return (self.total_avail + evictable_count) >= num_blocks def allocate_memory_blocks(self, num_blocks: int) -> Optional[Tensor]: @@ -205,11 +231,20 @@ def release_memory_blocks(self, blocks: Tensor) -> None: return if self.enable_prefix_caching: - self.block_ref_counts[blocks] -= 1 + # When multiple requests that share the same prefix finish on the same step, + # their block IDs appear multiple times in the blocks tensor. + # Writing `self.block_ref_counts[blocks] -= 1` would only decrement reference counts + # once per unique block. This is wrong. The reference counts must be decremented + # once per occurrence of the block in the `blocks` tensor. We need `scatter`. + blocks_i64 = blocks.to(torch.int64) + self.block_ref_counts.scatter_add_( + 0, blocks_i64, torch.full_like(blocks_i64, -1, dtype=torch.int32) + ) if self.prefix_caching_eviction_policy == PrefixCachingEvictionPolicy.REF_ZERO: zero_mask = self.block_ref_counts[blocks] == 0 if zero_mask.any(): - self._deregister_blocks(blocks[zero_mask]) + # Deduplicate so a shared block is deregistered/returned once. + self._deregister_blocks(torch.unique(blocks[zero_mask])) elif self.prefix_caching_eviction_policy == PrefixCachingEvictionPolicy.LRU: # Unregistered blocks (hash == -1, ref_count == 0) have no hash # entry to preserve for reuse (e.g., partial blocks at the end of @@ -219,7 +254,8 @@ def release_memory_blocks(self, blocks: Tensor) -> None: self.block_hashes[blocks] == -1 ) if unreg_mask.any(): - unreg_blocks = blocks[unreg_mask] + # Deduplicate so a shared block returns to the pool once. + unreg_blocks = torch.unique(blocks[unreg_mask]) num_unreg = unreg_blocks.numel() self.block_bag[self.total_avail : self.total_avail + num_unreg] = unreg_blocks self.total_avail += num_unreg @@ -256,6 +292,8 @@ def reset(self) -> None: self.block_ref_counts.fill_(0) if self.prefix_caching_eviction_policy == PrefixCachingEvictionPolicy.LRU: self.block_timestamps.fill_(0) + self.block_parent_id.fill_(-1) + self.block_child_count.fill_(0) # Clear per-block routing storage self.block_routing.clear() @@ -264,20 +302,54 @@ def reset(self) -> None: # Prefix caching methods # ========================================================================= - def register_kv_block_hashes(self, block_ids: list[int], block_hashes: list[int]) -> None: + def register_kv_block_hashes( + self, + block_ids: list[int], + block_hashes: list[int], + parent_hashes: Optional[list[int]] = None, + ) -> None: """Register blocks in the hash-to-block mapping for discovery (batch). Args: block_ids: List of block IDs. block_hashes: List of computed hash values (same length as block_ids). + parent_hashes: Parent hash for each block in the prefix chain (same + length as block_ids); 0 marks a root block with no parent. Used + by LRU eviction to avoid evicting a parent before its children. + If None, parents default to 0. """ if not block_ids: return id_tensor = torch.tensor(block_ids, dtype=torch.int64, device=self.block_hashes.device) hash_tensor = torch.tensor(block_hashes, dtype=torch.int64, device=self.block_hashes.device) self.block_hashes[id_tensor] = hash_tensor + if parent_hashes is not None: + assert len(parent_hashes) == len(block_ids) + # Add the new blocks to the hash map first so that a block whose parent is + # elsewhere in this same batch (block k's parent is block k-1) resolves. self.kv_hash_to_block_id.update(zip(block_hashes, block_ids)) + if self.prefix_caching_eviction_policy == PrefixCachingEvictionPolicy.LRU: + # Persist the resolved parent block id and bump each parent's child count. + # Parents are earlier in the prefix chain and already registered + # (a matched block or a prior chunk / earlier entry in this batch), + # so a valid parent hash resolves; 0 marks a root and an unknown hash + # falls back to -1. + if parent_hashes is None: + parent_hashes = [0] * len(block_ids) + parent_ids = [ + self.kv_hash_to_block_id.get(ph, -1) if ph != 0 else -1 for ph in parent_hashes + ] + parent_id_tensor = torch.tensor(parent_ids, dtype=torch.int64, device=id_tensor.device) + self.block_parent_id[id_tensor] = parent_id_tensor + has_parent = parent_id_tensor >= 0 + if has_parent.any(): + self.block_child_count.scatter_add_( + 0, + parent_id_tensor[has_parent], + torch.ones(int(has_parent.sum()), dtype=torch.int64), + ) + def _deregister_blocks(self, block_ids: Tensor) -> None: """Remove blocks from prefix caching state and return to free pool. @@ -306,10 +378,23 @@ def _deregister_blocks(self, block_ids: Tensor) -> None: self.on_blocks_deregistered(block_ids.tolist(), keys_to_delete) # Reset block state (batched tensor ops) - self.block_hashes[block_ids] = -1 - self.block_ref_counts[block_ids] = 0 if self.prefix_caching_eviction_policy == PrefixCachingEvictionPolicy.LRU: + # Drop these blocks from their parents' child counts before clearing + # their own bookkeeping, keeping block_child_count in sync so a parent + # becomes an evictable leaf once its last child is deregistered. + parent_ids = self.block_parent_id[block_ids_i64] + has_parent = parent_ids >= 0 + if has_parent.any(): + self.block_child_count.scatter_add_( + 0, + parent_ids[has_parent], + torch.full((int(has_parent.sum()),), -1, dtype=torch.int64), + ) + self.block_parent_id[block_ids] = -1 + self.block_child_count[block_ids] = 0 self.block_timestamps[block_ids] = 0 + self.block_hashes[block_ids] = -1 + self.block_ref_counts[block_ids] = 0 # Return blocks to free pool self.block_bag[self.total_avail : self.total_avail + num_blocks] = block_ids @@ -340,7 +425,53 @@ def get_evictable_block_count(self) -> Tensor: def evict_lru_blocks(self, num_blocks_needed: int) -> bool: """Evict LRU cached blocks to free up space in the pool. - Evicts blocks with ref_count == 0, starting with oldest timestamps. + Evicts blocks with ref_count == 0, least-recently-used first, while never + evicting a parent before its children. Block hashes are parent-chained, + and ``_find_kv_match_count`` relies on the invariant that a cached child + block always has all of its ancestors cached too. A naive oldest-first + eviction breaks this: with chunked prefill, earlier chunks are allocated + first (older timestamps) yet are ancestors of later chunks (newer + timestamps), so once the request finishes and its blocks are cached, an + ancestor can be older than its descendant and get evicted first, leaving a + dangling child. + + To preserve the invariant while staying optimal we peel the cached forest + from its leaves inward with a min-heap: only a leaf (a cached block with + no cached children) is ever evictable, and among the currently-evictable + leaves we always take the one with the oldest *own* timestamp. Evicting a + leaf can turn its parent into a leaf, which is then pushed onto the heap. + Repeating ``num_blocks_needed`` times gives, at each step, the globally + least-recently-used block that can be removed without orphaning a child — + the natural generalization of LRU to the parent-chain constraint. Keying + each block by its *own* recency (and only reconsidering a parent once its + children are gone) is what makes this optimal: a block is retained purely + because it is recently used, never because a hot descendant props it up, + so a colder evictable block is always evicted before a hotter one. + + Worked example, evicting 3 from:: + + A(ts 1) -> B(ts 2) -> C(ts 5) (C, F are leaves under B) + \-> F(ts 3) + \-> D(ts 3) -> E(ts 5) (E is a leaf under D) + + Leaf-peel evicts F(3), then C(5); B is now childless so it joins the + leaves with its own ts=2 and is evicted next -> retains {A, D, E}, keeping + the hottest block E(5) rather than the colder interior block B(2). + + Note: because a request holds a contiguous block prefix [0..k], any in-use + (ref_count > 0) block keeps all of its ancestors in use too. Hence a cached + (ref_count == 0) block can only have cached children, and considering the + cached set alone is sufficient to avoid dangling children. + + The parent block id of each block and its live child count are maintained + incrementally on register/deregister (``block_parent_id`` / + ``block_child_count``), so this method reads the prefix forest directly + rather than rebuilding it from hashes with a per-eviction sort. Only the + inherently-sequential leaf peel below is per-element. + + The parent graph is assumed acyclic (a forest), which holds for any hashes + produced by the prefix-chain builder; an assertion guards against a + pathological hash collision wedging the peel. Args: num_blocks_needed: Number of blocks to evict. @@ -352,14 +483,49 @@ def evict_lru_blocks(self, num_blocks_needed: int) -> bool: cached_mask = (self.block_ref_counts == 0) & (self.block_hashes != -1) cached_block_ids = torch.nonzero(cached_mask, as_tuple=True)[0] - if cached_block_ids.numel() < num_blocks_needed: + num_cached = cached_block_ids.numel() + if num_cached < num_blocks_needed: return False # Not enough cached blocks to evict + if num_blocks_needed <= 0: + return True - # Sort by timestamp (ascending = oldest first) - cached_timestamps = self.block_timestamps[cached_block_ids] - sorted_indices = torch.argsort(cached_timestamps) - blocks_to_evict = cached_block_ids[sorted_indices[:num_blocks_needed]] + ts = self.block_timestamps[cached_block_ids].tolist() + bid = cached_block_ids.tolist() + parent_global = self.block_parent_id[cached_block_ids].tolist() + child_count = self.block_child_count[cached_block_ids].tolist() + + # Map a cached block's global id to its local index so the peel can find a + # parent's slot to decrement. Parents that are not cached (root, or a + # parent still in use) are absent and are simply treated as peel roots. + global_to_local = {bid[i]: i for i in range(num_cached)} + parent_local = [global_to_local.get(p, -1) for p in parent_global] + + # Min-heap of currently-evictable leaves keyed by (own timestamp, block + # id). Block ids are unique, so the tie-break is total and deterministic. + heap = [(ts[i], bid[i], i) for i in range(num_cached) if child_count[i] == 0] + heapq.heapify(heap) + + evicted_local = [] + while heap and len(evicted_local) < num_blocks_needed: + _, _, i = heapq.heappop(heap) + evicted_local.append(i) + p = parent_local[i] + if p >= 0: + child_count[p] -= 1 + if child_count[p] == 0: + heapq.heappush(heap, (ts[p], bid[p], p)) + + # A forest is always fully peelable, so the heap always exposes enough + # leaves to collect num_blocks_needed (guaranteed by the num_cached >= + # num_blocks_needed check above). Falling short means the parent graph is + # cyclic — only possible under a hash collision, which we treat as a bug. + assert len(evicted_local) == num_blocks_needed, ( + f"leaf peel evicted {len(evicted_local)} of {num_blocks_needed} " + f"requested from {num_cached} cached blocks; parent graph is not a " + f"forest (likely a block-hash collision)" + ) + blocks_to_evict = cached_block_ids[torch.tensor(evicted_local, dtype=torch.int64)] self._deregister_blocks(blocks_to_evict) return True diff --git a/megatron/core/inference/contexts/mamba_slot_allocator.py b/megatron/core/inference/contexts/mamba_slot_allocator.py index 60c8dd3416b..e7b977a4f30 100644 --- a/megatron/core/inference/contexts/mamba_slot_allocator.py +++ b/megatron/core/inference/contexts/mamba_slot_allocator.py @@ -100,8 +100,13 @@ def __init__( # CPU flag to skip GPU sync when no intermediates exist self._has_intermediates = False - # Pre-allocated output buffers for CUDA graph compatible extraction (GPU). - self.max_intermediate_count = MAX_INTERMEDIATE_OFFSETS_PER_REQUEST * context.max_requests + # Pre-allocated "scratch" output buffers for CUDA graph compatible + # extraction (GPU): per-step staging that the kernel writes intermediate + # states into before commit copies them to the durable cache above. Sized + # by the per-step token budget computed once on the context; the budget + # accounting in DynamicInferenceContext refers to these as the "scratch" + # buffers. + self.max_intermediate_count = context.max_mamba_intermediate_states_per_step self.intermediate_ssm_out = torch.zeros( (num_mamba_layers, self.max_intermediate_count) + ssm_states_shape, dtype=ssm_states_dtype, @@ -402,24 +407,38 @@ def compute_and_store_offsets( overall_required_blocks: Total blocks needed for this request. """ ctx = self.context + bs = ctx.block_size_tokens prompt_len = len(req.prompt_tokens) - num_kv_matched = num_matched_blocks - kv_div_abs = num_kv_matched * ctx.block_size_tokens - last_aligned_abs = (prompt_len // ctx.block_size_tokens) * ctx.block_size_tokens - seq_len = prefill_chunk_length - skip_tokens # effective prefill length - # Compute relative offsets (relative to prefill start after skip) - kv_div_rel = kv_div_abs - skip_tokens - last_aligned_rel = last_aligned_abs - skip_tokens - penultimate_abs = (overall_required_blocks - 1) * ctx.block_size_tokens - penultimate_rel = penultimate_abs - skip_tokens - - # Determine mamba_chunk_size from mamba config (128 is the standard SSM kernel chunk size) - mamba_chunk_size = 128 - - # Build offset list: include if > 0, < seq_len, and % mamba_chunk_size == 0 + # Absolute token position (from the prompt start) where THIS chunk's + # computed tokens begin. The first chunk computes from `skip_tokens` (the + # prefix that was skipped); continuation chunks compute from + # `finished_chunk_token_count` (with skip_tokens == 0). Framing the + # boundary offsets against this chunk start -- rather than assuming the + # first chunk -- lets us extract Mamba state at block boundaries that fall + # in ANY chunk. In particular the last complete block of a multi-chunk + # prompt lives in a continuation chunk; it was previously unreachable, so + # non-block-aligned prompts never cached a usable resume boundary and + # later turns could not skip prefill. + chunk_start = req.finished_chunk_token_count + skip_tokens + seq_len = prefill_chunk_length - skip_tokens # tokens computed this chunk + is_last_chunk = req.finished_chunk_token_count + prefill_chunk_length >= prompt_len + + # Candidate absolute block boundaries at which to cache Mamba state. + kv_div_abs = num_matched_blocks * bs + last_aligned_abs = (prompt_len // bs) * bs # last complete block boundary + penultimate_abs = (overall_required_blocks - 1) * bs + + # SSM chunk size the mamba kernel actually runs with. States can only be + # extracted at multiples of this value, and it must match the value used + # in MambaMetadata (offset -> chunk-index conversion) to stay consistent. + mamba_chunk_size = ctx.mamba_chunk_size + + # Keep only boundaries that land inside this chunk's computed tokens and on + # a mamba-chunk boundary (required for mid-sequence state extraction). offsets_set = set() - for offset in [kv_div_rel, last_aligned_rel, penultimate_rel]: + for abs_pos in (kv_div_abs, last_aligned_abs, penultimate_abs): + offset = abs_pos - chunk_start if offset > 0 and offset < seq_len and offset % mamba_chunk_size == 0: offsets_set.add(offset) @@ -428,8 +447,8 @@ def compute_and_store_offsets( # CPU bookkeeping writes (no GPU kernel launches). if count > 0: - abs_tokens_cpu = torch.tensor([skip_tokens + o for o in offsets], dtype=torch.int64) - block_indices_cpu = abs_tokens_cpu // ctx.block_size_tokens - 1 + abs_tokens_cpu = torch.tensor([chunk_start + o for o in offsets], dtype=torch.int64) + block_indices_cpu = abs_tokens_cpu // bs - 1 bids_cpu = ctx.request_to_kv_block_ids[current_id][block_indices_cpu] self._intermediate_offsets_cpu[current_id, :count] = torch.tensor( @@ -439,9 +458,13 @@ def compute_and_store_offsets( self._has_intermediates = True self._intermediate_counts_cpu[current_id] = count - # Block-aligned EOS: prompt_len is exactly block-aligned - if last_aligned_abs == prompt_len and prompt_len > 0: - last_block_idx = prompt_len // ctx.block_size_tokens - 1 + # Block-aligned EOS: when the prompt length is exactly block-aligned, the + # request's live final state IS the last block boundary's state and can be + # cached directly. Only valid on the final chunk (otherwise the live state + # is mid-prompt). Non-block-aligned prompts cache their last complete block + # via the intermediate-extraction path above instead. + if is_last_chunk and last_aligned_abs == prompt_len and prompt_len > 0: + last_block_idx = prompt_len // bs - 1 if last_block_idx >= 0: self._eos_cache_block_id_cpu[current_id] = ctx.request_to_kv_block_ids[current_id][ last_block_idx diff --git a/megatron/core/inference/data_parallel_inference_coordinator/__init__.py b/megatron/core/inference/data_parallel_inference_coordinator/__init__.py new file mode 100644 index 00000000000..097b9f28f72 --- /dev/null +++ b/megatron/core/inference/data_parallel_inference_coordinator/__init__.py @@ -0,0 +1,22 @@ +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Data parallel inference coordinator package. + +The coordinator class itself lives in coordinator.py; message handlers are in +handlers.py and the control-signal state machine in state.py. This module +re-exports the public names so existing imports of +``megatron.core.inference.data_parallel_inference_coordinator`` keep working. +""" + +from .coordinator import DataParallelInferenceCoordinator +from .handlers import HANDLERS, message_handler +from .state import CONTROL_TRANSITIONS, ControlTransition, CoordinatorState + +__all__ = [ + "DataParallelInferenceCoordinator", + "CoordinatorState", + "ControlTransition", + "CONTROL_TRANSITIONS", + "HANDLERS", + "message_handler", +] diff --git a/megatron/core/inference/data_parallel_inference_coordinator.py b/megatron/core/inference/data_parallel_inference_coordinator/coordinator.py similarity index 64% rename from megatron/core/inference/data_parallel_inference_coordinator.py rename to megatron/core/inference/data_parallel_inference_coordinator/coordinator.py index 50f586cc598..6900887bcf0 100644 --- a/megatron/core/inference/data_parallel_inference_coordinator.py +++ b/megatron/core/inference/data_parallel_inference_coordinator/coordinator.py @@ -7,7 +7,6 @@ import signal import socket from collections import deque -from enum import Enum, auto from multiprocessing import Event from multiprocessing.connection import Connection @@ -21,6 +20,9 @@ TextGenerationController, ) +from .handlers import HANDLERS +from .state import CoordinatorState + try: import zmq @@ -55,19 +57,23 @@ class DataParallelInferenceCoordinator: `InferenceClient`, and performs a simple handshake. 3. **Request Forwarding**: It receives inference requests from clients, assigns a unique server-side request ID, tokenizes the prompt, and forwards the request - to one of the available data parallel rank using a round-robin scheduling - strategy. + to one of the available data parallel ranks using load-balanced (and, + when prefix caching is enabled, prefix-affinity-aware) routing. 4. **Response Routing**: It receives completed results from the data parallel ranks and routes them back to the original client that made the request. 5. **Control Signal Broadcasting**: It relays control signals (e.g., PAUSE, STOP) from a client to all connected data parallel ranks. + Message handling is split out into handlers.py: the event loop in start() + dispatches each message to the handler registered for its header, so + supporting a new message type requires no changes here. + Attributes: router_socket (zmq.Socket): The central ZMQ ROUTER socket for all communication. data_parallel_size (int): The number of data parallel workers to expect. identities_of_data_parallel_ranks (deque): A deque holding the ZMQ - identities of connected TP-coordinators, used for round-robin scheduling. + identities of connected data parallel instances, used for request routing. request_id_to_client_id (dict): Maps server-side request IDs to the ZMQ identity of the client that initiated the request. request_id_to_client_request_id (dict): Maps server-side request IDs to the @@ -75,13 +81,9 @@ class DataParallelInferenceCoordinator: next_request_id (int): A counter for generating unique server-side request IDs. """ - class CoordinatorState(Enum): - """State machine for the coordinator.""" - - RUNNING = auto() - PAUSED = auto() - SUSPENDED = auto() - STOPPING = auto() + # Exposed as a class attribute for backwards compatibility; the canonical + # definition lives in state.py. + CoordinatorState = CoordinatorState def __init__( self, @@ -109,7 +111,7 @@ def __init__( Args: pipe_connection (Connection): A connecting pipe to the parent process. - data_parallel_size (int): The number of TP-coordinator workers that are + data_parallel_size (int): The number of data parallel instances that are expected to connect. tokenizer: The tokenizer to use for prompt tokenization and detokenization. inference_coordinator_port (Optional[int]): The TCP port number to bind the server to. @@ -182,15 +184,15 @@ def __init__( self.identities_of_data_parallel_ranks = deque( sorted(self.identities_of_data_parallel_ranks) ) - self._round_robin_idx = 0 self.request_id_to_client_id = {} self.request_id_to_client_request_id = {} self.request_id_to_rank = {} # Maps request_id → rank identity for pending count tracking + self.removed_engine_identities = set() self.next_request_id = 0 self.tokenizer = tokenizer - self.state = self.CoordinatorState.RUNNING + self.state = CoordinatorState.RUNNING # Prefix caching state for routing. self.block_size_tokens = block_size_tokens @@ -221,19 +223,25 @@ def __init__( self._hash_table: dict[int, dict[int, int]] = {} self._hash_assignment_counter = 0 - def get_next_data_parallel_rank(self): + # Clients that have completed the CONNECT handshake. + self.known_clients = set() + + # Header -> handler dispatch table, sourced from the handler registry. + self._handlers = dict(HANDLERS) + + def get_least_loaded_data_parallel_rank(self): """ - Selects the next data parallel rank using round-robin scheduling. + Selects the data parallel rank with the fewest in-flight requests. + + Ties are broken by lowest rank index for deterministic behavior. Returns: - bytes: The ZMQ identity of the next data parallel rank to receive a request. + bytes: The ZMQ identity of the least-loaded data parallel rank. """ - identities = self.identities_of_data_parallel_ranks - if not identities: + if not self._identities_list: raise RuntimeError("No engines connected") - idx = self._round_robin_idx % len(identities) - self._round_robin_idx = idx + 1 - return identities[idx] + best_idx = int(np.argmin(self._pending_counts)) + return self._identities_list[best_idx] def _register_rank_identity(self, identity): """Register a new rank identity in the scoring data structures. @@ -256,8 +264,38 @@ def _register_rank_identity(self, identity): ) def _remove_engine(self, identity): - """Remove a disconnected engine from the routing pool.""" + """Remove a disconnected engine from all routing bookkeeping. + Called both during shutdown and when an engine becomes unreachable mid-operation + (e.g. zmq.EHOSTUNREACH in _send_to_engine). The O(n) index-shifting and hash-table + rebuild are acceptable because the number of connected engines is small; optimize + only if dynamic registration/deregistration at high engine counts becomes a use case. + """ self.identities_of_data_parallel_ranks.remove(identity) + self.removed_engine_identities.add(identity) + idx = self.identity_to_rank_index.pop(identity, None) + if idx is None: + return + self._identities_list.pop(idx) + self._pending_counts = np.delete(self._pending_counts, idx) + # Shift indices for engines that came after the removed slot. + for ident in self.identity_to_rank_index: + if self.identity_to_rank_index[ident] > idx: + self.identity_to_rank_index[ident] -= 1 + # Drop hash-table entries for the removed rank; shift indices above it. + new_hash_table = {} + for h, rank_ts in self._hash_table.items(): + # h is hash index + # rank_ts is a dict mapping rank_idx → timestamp + new_row = {} + for r, ts in rank_ts.items(): + if r == idx: + # skip this rank as it is removed + continue + new_r = r - 1 if r > idx else r + new_row[new_r] = ts + if new_row: + new_hash_table[h] = new_row + self._hash_table = new_hash_table logging.warning( "Coordinator: removed engine %s (now %d engines)", identity, @@ -279,6 +317,12 @@ def _send_to_engine(self, identity, payload): return False raise + def _broadcast_to_engines(self, payload): + """Send a deserialized payload to every connected data parallel rank.""" + serialized = msgpack.packb(payload, use_bin_type=True) + for data_parallel_rank_id in list(self.identities_of_data_parallel_ranks): + self._send_to_engine(data_parallel_rank_id, serialized) + def compute_request_hashes(self, prompt): """Compute block hashes for a prompt on CPU. @@ -312,11 +356,13 @@ def get_best_data_parallel_rank(self, request_hashes): Returns: bytes: The ZMQ identity of the selected data parallel rank. """ - if self.prefix_caching_coordinator_policy == PrefixCachingCoordinatorPolicy.ROUND_ROBIN: - return self.get_next_data_parallel_rank() + if self.prefix_caching_coordinator_policy == PrefixCachingCoordinatorPolicy.LOAD_BALANCED: + return self.get_least_loaded_data_parallel_rank() + # Without prefix caching (or when the request has no hashes to match on) + # fall back to load-balanced routing. if not self.enable_prefix_caching or not request_hashes: - return self.get_next_data_parallel_rank() + return self.get_least_loaded_data_parallel_rank() match, recency = self._match_vector(request_hashes) @@ -378,202 +424,33 @@ def start(self): Starts the main event loop for the coordinator. This method runs an infinite loop, continuously listening for incoming - messages on the ZMQ ROUTER socket. It parses the message header to - determine the message type and takes appropriate action, such as - handling new client connections, forwarding requests, broadcasting - control signals, or processing replies from the engines. + messages on the ZMQ ROUTER socket. It reads the message header and + dispatches to the handler registered for it (see handlers.py). + A handler that returns a truthy value stops the loop. """ # Todo [Siddharth]: Make this more robust to handle invalid messages. - known_clients = set() while True: sender_identity, serialized_payload = self.router_socket.recv_multipart() - # Allow for re-registration if connecting to a running coordinator. + # An empty payload is a data parallel rank (re-)registering itself. if serialized_payload == b"": - if sender_identity not in self.identities_of_data_parallel_ranks: - self.identities_of_data_parallel_ranks.append(sender_identity) - self._register_rank_identity(sender_identity) + self._handle_rank_registration(sender_identity) continue deserialized_payload = msgpack.unpackb(serialized_payload, raw=False) header = Headers(deserialized_payload[0]) - if header == Headers.CONNECT: - if sender_identity in known_clients: - logging.info( - f"Client {sender_identity} sent a duplicate connect request. Ignoring .." - ) - continue - - # print(f"New client connected: {sender_identity}") - known_clients.add(sender_identity) - self.router_socket.send_multipart( - [sender_identity, msgpack.packb([Headers.CONNECT_ACK.value], use_bin_type=True)] - ) - - elif header == Headers.SUBMIT_REQUEST: - # ToDo [Siddharth]: We might want to tokenize the prompt on the - # assigned data parallel rank for this process instead - # of the coordinator. - - # Message from a known client - if sender_identity not in known_clients: - logging.info( - f"Received message from unknown client {sender_identity}. Ignoring." - ) - continue - # this is a message from a client. - # route it to a data parallel rank - client_request_id, prompt, sampling_params = deserialized_payload[1:] - # map client request_id to server request_id - # necessary because multiple clients might have the same request_id. - request_id = self.next_request_id - self.next_request_id += 1 - self.request_id_to_client_id[request_id] = sender_identity - self.request_id_to_client_request_id[request_id] = client_request_id - - # Serialize prompt. - if isinstance(prompt, (str, list)): - pass - elif isinstance(prompt, torch.Tensor): - prompt = prompt.tolist() - else: - raise Exception("specialize for <%s> prompt." % type(prompt).__name__) - - payload = msgpack.packb( - [Headers.SUBMIT_REQUEST.value, request_id, prompt, sampling_params], - use_bin_type=True, - ) - - request_hashes = self.compute_request_hashes(prompt) - if ( - self.prefix_caching_coordinator_policy - == PrefixCachingCoordinatorPolicy.FIRST_PREFIX_BLOCK - ): - request_hashes = request_hashes[:1] - - # Account for the fact that some engines may have died. - for _ in range(len(self.identities_of_data_parallel_ranks)): - next_identity = self.get_best_data_parallel_rank(request_hashes) - if self._send_to_engine(next_identity, payload): - break - else: - # If all engines have died, we are in an abnormal state, and must exit cleanly. - logging.error("Coordinator: no reachable engines for request %d", request_id) - del self.request_id_to_client_id[request_id] - del self.request_id_to_client_request_id[request_id] - return - - self.request_id_to_rank[request_id] = next_identity - self._pending_counts[self.identity_to_rank_index[next_identity]] += 1 - if request_hashes: - self._update_rank_hashes(next_identity, request_hashes) - if self.schedule_records is not None: - self.schedule_records.append( - { - "request_id": request_id, - "rank_index": self.identity_to_rank_index[next_identity], - "num_hashes": len(request_hashes), - } - ) - - elif header in ( - Headers.PAUSE, - Headers.UNPAUSE, - Headers.SUSPEND, - Headers.RESUME, - Headers.SET_GENERATION_EPOCH, - Headers.STOP, - ): - # Start by checking the current state against the control signal. - if sender_identity not in known_clients: - logging.warning("Coordinator: ignoring signal from unknown client.") - continue - - if header == Headers.PAUSE: - idem_states = (self.CoordinatorState.PAUSED, self.CoordinatorState.SUSPENDED) - if self.state == self.CoordinatorState.RUNNING: - self.state = self.CoordinatorState.PAUSED - elif self.state in idem_states: - # Already paused/suspended, ignore redundant PAUSE. - continue - else: - logging.warning("Coordinator: ignoring PAUSE in state %s", self.state) - continue - elif header == Headers.UNPAUSE: - if self.state != self.CoordinatorState.PAUSED: - logging.warning("Coordinator: ignoring UNPAUSE in state %s", self.state) - continue - self.state = self.CoordinatorState.RUNNING - elif header == Headers.SUSPEND: - if self.state != self.CoordinatorState.PAUSED: - logging.warning("Coordinator: ignoring SUSPEND in state %s", self.state) - continue - self.state = self.CoordinatorState.SUSPENDED - elif header == Headers.RESUME: - if self.state != self.CoordinatorState.SUSPENDED: - logging.warning("Coordinator: ignoring RESUME in state %s", self.state) - continue - self.state = self.CoordinatorState.PAUSED - elif header == Headers.STOP: - good_states = (self.CoordinatorState.PAUSED, self.CoordinatorState.SUSPENDED) - if self.state not in good_states: - logging.warning("Coordinator: ignoring STOP in state %s", self.state) - continue - self.state = self.CoordinatorState.STOPPING - - # Broadcast the control signal if we're in a good state. - # Forward the full deserialized payload so that data-bearing - # signals (e.g. SET_GENERATION_EPOCH) retain their arguments. - broadcast_payload = msgpack.packb(deserialized_payload, use_bin_type=True) - for data_parallel_rank_id in list(self.identities_of_data_parallel_ranks): - self._send_to_engine(data_parallel_rank_id, broadcast_payload) - - # STOP affects engines; reset coordinator to RUNNING to allow future engines. - if header == Headers.STOP: - self.state = self.CoordinatorState.RUNNING - - elif header == Headers.ENGINE_REPLY: - # This is the output of a single engine step on some data parallel rank. - assert sender_identity in self.identities_of_data_parallel_ranks - finished_requests = deserialized_payload[1] - - for finished_request in finished_requests: - self.detokenize(finished_request) - fid = finished_request["request_id"] - client_identity = self.request_id_to_client_id[fid] - client_request_identity = self.request_id_to_client_request_id[fid] - del self.request_id_to_client_id[fid] - del self.request_id_to_client_request_id[fid] - assigned_rank = self.request_id_to_rank.pop(fid, None) - if assigned_rank is not None: - idx = self.identity_to_rank_index.get(assigned_rank) - if idx is not None: - assert self._pending_counts[idx] >= 1 - self._pending_counts[idx] -= 1 - - self.router_socket.send_multipart( - [ - client_identity, - msgpack.packb( - [header.value, client_request_identity, finished_request], - use_bin_type=True, - ), - ] - ) - - elif header == Headers.SHUTDOWN: - if sender_identity not in known_clients: - logging.warning("Coordinator: ignoring signal from unknown client.") - continue + handler = self._handlers.get(header) + if handler is None: + raise UnknownHeaderError(header) + if handler(self, sender_identity, deserialized_payload): break - elif header == Headers.DISCONNECT: - if sender_identity in self.identities_of_data_parallel_ranks: - self._remove_engine(sender_identity) - - else: - raise UnknownHeaderError(header) + def _handle_rank_registration(self, sender_identity): + """Register a data parallel rank that connected to a running coordinator.""" + if sender_identity not in self.identities_of_data_parallel_ranks: + self.identities_of_data_parallel_ranks.append(sender_identity) + self._register_rank_identity(sender_identity) def detokenize(self, finished_request): """ @@ -586,10 +463,6 @@ def detokenize(self, finished_request): finished_request (dict): The serialized merged request containing the generated tokens to be detokenized. It is modified in place. """ - if finished_request["prompt"] is None: - finished_request["prompt"] = TextGenerationController.detokenize( - self.tokenizer, finished_request["prompt_tokens"][1], remove_EOD=False - ) detokenize_stop_sequence = (finished_request.get("sampling_params", {}) or {}).get( "detokenize_stop_sequence", False ) @@ -629,7 +502,7 @@ def entrypoint( ready_event (Event): A threading or multiprocessing event object that is set() once the coordinator is ready to accept connections. inference_coordinator_port (int): The port to bind to. - data_parallel_size (int): The number of expected TP-coordinators. + data_parallel_size (int): The number of expected data parallel instances. deterministic_mode (bool): Whether to enable deterministic scheduling. block_size_tokens (Optional[int]): Token block size for prefix caching hashing. enable_prefix_caching (bool): Whether prefix caching is enabled. diff --git a/megatron/core/inference/data_parallel_inference_coordinator/handlers.py b/megatron/core/inference/data_parallel_inference_coordinator/handlers.py new file mode 100644 index 00000000000..b2d9a36fffb --- /dev/null +++ b/megatron/core/inference/data_parallel_inference_coordinator/handlers.py @@ -0,0 +1,233 @@ +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Message handlers for the data parallel inference coordinator. + +Each handler is a free function decorated with @message_handler, which records +it in the module-level HANDLERS registry keyed by message header. The +coordinator builds its dispatch table from this registry, so a new message type +is supported simply by adding a decorated function here; the coordinator's event +loop never changes. + +Handlers have the signature ``(coordinator, sender_identity, payload) -> bool | None`` +where ``payload`` is the already-deserialized message. Returning a truthy value +signals the coordinator's event loop to stop. +""" + +import logging + +import torch + +from megatron.core.inference.config import PrefixCachingCoordinatorPolicy +from megatron.core.inference.headers import Headers + +from .state import CONTROL_TRANSITIONS, CoordinatorState + +try: + import msgpack +except ImportError: + msgpack = None + + +# Maps a message header value to the function that handles it. Populated by the +# @message_handler decorator at import time. +HANDLERS = {} + + +def message_handler(*headers): + """Register a function as the handler for one or more message headers. + + A new message type is supported by writing a handler function and decorating + it with the header(s) it serves; it is added to HANDLERS, which the + coordinator turns into its dispatch table. The event loop never needs to + change when a header is added. + """ + + def decorator(fn): + for header in headers: + assert header not in HANDLERS, f"duplicate handler for {header}" + HANDLERS[header] = fn + return fn + + return decorator + + +@message_handler(Headers.CONNECT) +def handle_connect(coordinator, sender_identity, payload): + """Handshake with a new client, replying with a CONNECT_ACK.""" + if sender_identity in coordinator.known_clients: + logging.info(f"Client {sender_identity} sent a duplicate connect request. Ignoring ..") + return + + coordinator.known_clients.add(sender_identity) + coordinator.router_socket.send_multipart( + [sender_identity, msgpack.packb([Headers.CONNECT_ACK.value], use_bin_type=True)] + ) + + +@message_handler(Headers.SUBMIT_REQUEST) +def handle_submit_request(coordinator, sender_identity, payload): + """Route a client request to a data parallel rank. + + Returns True (stopping the loop) if no engines are reachable. + """ + # ToDo [Siddharth]: We might want to tokenize the prompt on the + # assigned data parallel rank for this process instead + # of the coordinator. + + # Message from a known client + if sender_identity not in coordinator.known_clients: + logging.info(f"Received message from unknown client {sender_identity}. Ignoring.") + return + # this is a message from a client. + # route it to a data parallel rank + client_request_id, prompt, sampling_params = payload[1:] + # map client request_id to server request_id + # necessary because multiple clients might have the same request_id. + request_id = coordinator.next_request_id + coordinator.next_request_id += 1 + coordinator.request_id_to_client_id[request_id] = sender_identity + coordinator.request_id_to_client_request_id[request_id] = client_request_id + + # Serialize prompt. + if isinstance(prompt, (str, list)): + pass + elif isinstance(prompt, torch.Tensor): + prompt = prompt.tolist() + else: + raise Exception("specialize for <%s> prompt." % type(prompt).__name__) + + engine_payload = msgpack.packb( + [Headers.SUBMIT_REQUEST.value, request_id, prompt, sampling_params], use_bin_type=True + ) + + request_hashes = coordinator.compute_request_hashes(prompt) + if ( + coordinator.prefix_caching_coordinator_policy + == PrefixCachingCoordinatorPolicy.FIRST_PREFIX_BLOCK + ): + request_hashes = request_hashes[:1] + + # Account for the fact that some engines may have died. + for _ in range(len(coordinator.identities_of_data_parallel_ranks)): + next_identity = coordinator.get_best_data_parallel_rank(request_hashes) + if coordinator._send_to_engine(next_identity, engine_payload): + break + else: + # If all engines have died, we are in an abnormal state, and must exit cleanly. + logging.error("Coordinator: no reachable engines for request %d", request_id) + del coordinator.request_id_to_client_id[request_id] + del coordinator.request_id_to_client_request_id[request_id] + return True + + coordinator.request_id_to_rank[request_id] = next_identity + coordinator._pending_counts[coordinator.identity_to_rank_index[next_identity]] += 1 + if request_hashes: + coordinator._update_rank_hashes(next_identity, request_hashes) + if coordinator.schedule_records is not None: + coordinator.schedule_records.append( + { + "request_id": request_id, + "rank_index": coordinator.identity_to_rank_index[next_identity], + "num_hashes": len(request_hashes), + } + ) + + +@message_handler( + Headers.PAUSE, + Headers.UNPAUSE, + Headers.SUSPEND, + Headers.RESUME, + Headers.SET_GENERATION_EPOCH, + Headers.STOP, +) +def handle_control_signal(coordinator, sender_identity, payload): + """Validate a control signal against the transition table and broadcast it.""" + if sender_identity not in coordinator.known_clients: + logging.warning("Coordinator: ignoring signal from unknown client.") + return + + header = Headers(payload[0]) + transition = CONTROL_TRANSITIONS[header] + if coordinator.state not in transition.allowed_from: + # Silently ignore redundant signals; warn on genuinely invalid ones. + if coordinator.state not in transition.idempotent_in: + logging.warning("Coordinator: ignoring %s in state %s", header.name, coordinator.state) + return + if transition.new_state is not None: + coordinator.state = transition.new_state + + # Broadcast the control signal. Forward the full deserialized payload so + # that data-bearing signals (e.g. SET_GENERATION_EPOCH) retain their args. + coordinator._broadcast_to_engines(payload) + + # STOP affects engines; reset coordinator to RUNNING to allow future engines. + if header == Headers.STOP: + coordinator.state = CoordinatorState.RUNNING + + +@message_handler(Headers.START_CUDA_PROFILER, Headers.STOP_CUDA_PROFILER) +def handle_cuda_profiler_signal(coordinator, sender_identity, payload): + """Broadcast a CUDA profiler control signal to every connected DP engine. + + Profiler control is not a coordinator state transition, so there are no + CoordinatorState checks — the signal is simply forwarded to all engines. + """ + if sender_identity not in coordinator.known_clients: + logging.warning("Coordinator: ignoring profiler signal from unknown client.") + return + coordinator._broadcast_to_engines(payload) + + +@message_handler(Headers.ENGINE_REPLY) +def handle_engine_reply(coordinator, sender_identity, payload): + """Route completed requests from an engine back to their originating clients.""" + # This is the output of a single engine step on some data parallel rank. + if sender_identity not in coordinator.identities_of_data_parallel_ranks: + # A removed engine's final replies may still be queued up. + # Only exit with an assert if the sender was never connected to the coordinator. + assert ( + sender_identity in coordinator.removed_engine_identities + ), f"ENGINE_REPLY from never-connected sender {sender_identity!r}" + logging.warning("Coordinator: ENGINE_REPLY from removed engine %r", sender_identity) + finished_requests = payload[1] + + for finished_request in finished_requests: + coordinator.detokenize(finished_request) + fid = finished_request["request_id"] + client_identity = coordinator.request_id_to_client_id[fid] + client_request_identity = coordinator.request_id_to_client_request_id[fid] + del coordinator.request_id_to_client_id[fid] + del coordinator.request_id_to_client_request_id[fid] + assigned_rank = coordinator.request_id_to_rank.pop(fid, None) + if assigned_rank is not None: + idx = coordinator.identity_to_rank_index.get(assigned_rank) + if idx is not None: + assert coordinator._pending_counts[idx] >= 1 + coordinator._pending_counts[idx] -= 1 + + coordinator.router_socket.send_multipart( + [ + client_identity, + msgpack.packb( + [Headers.ENGINE_REPLY.value, client_request_identity, finished_request], + use_bin_type=True, + ), + ] + ) + + +@message_handler(Headers.SHUTDOWN) +def handle_shutdown(coordinator, sender_identity, payload): + """Stop the coordinator event loop on request from a known client.""" + if sender_identity not in coordinator.known_clients: + logging.warning("Coordinator: ignoring signal from unknown client.") + return + return True + + +@message_handler(Headers.DISCONNECT) +def handle_disconnect(coordinator, sender_identity, payload): + """Remove a disconnecting engine from the routing pool.""" + if sender_identity in coordinator.identities_of_data_parallel_ranks: + coordinator._remove_engine(sender_identity) diff --git a/megatron/core/inference/data_parallel_inference_coordinator/state.py b/megatron/core/inference/data_parallel_inference_coordinator/state.py new file mode 100644 index 00000000000..bb6c32104f8 --- /dev/null +++ b/megatron/core/inference/data_parallel_inference_coordinator/state.py @@ -0,0 +1,63 @@ +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""State machine definitions for the data parallel inference coordinator.""" + +from dataclasses import dataclass, field +from enum import Enum, auto + +from megatron.core.inference.headers import Headers + + +class CoordinatorState(Enum): + """State machine for the coordinator.""" + + RUNNING = auto() + PAUSED = auto() + SUSPENDED = auto() + STOPPING = auto() + + +_ALL_STATES = frozenset(CoordinatorState) + + +@dataclass(frozen=True) +class ControlTransition: + """A single rule in the control-signal state machine. + + Attributes: + allowed_from: States the signal may be applied from. + new_state: State to move to once applied, or None to leave the state + unchanged (e.g. a pure broadcast such as SET_GENERATION_EPOCH). + idempotent_in: States in which the signal is a silent no-op rather than + a logged rejection (e.g. a redundant PAUSE while already paused). + """ + + allowed_from: frozenset + new_state: CoordinatorState | None + idempotent_in: frozenset = field(default_factory=frozenset) + + +# Control-signal state machine, expressed declaratively and consumed by the +# control-signal handler. +CONTROL_TRANSITIONS = { + Headers.PAUSE: ControlTransition( + allowed_from=frozenset({CoordinatorState.RUNNING}), + new_state=CoordinatorState.PAUSED, + idempotent_in=frozenset({CoordinatorState.PAUSED, CoordinatorState.SUSPENDED}), + ), + Headers.UNPAUSE: ControlTransition( + allowed_from=frozenset({CoordinatorState.PAUSED}), new_state=CoordinatorState.RUNNING + ), + Headers.SUSPEND: ControlTransition( + allowed_from=frozenset({CoordinatorState.PAUSED}), new_state=CoordinatorState.SUSPENDED + ), + Headers.RESUME: ControlTransition( + allowed_from=frozenset({CoordinatorState.SUSPENDED}), new_state=CoordinatorState.PAUSED + ), + Headers.STOP: ControlTransition( + allowed_from=frozenset({CoordinatorState.PAUSED, CoordinatorState.SUSPENDED}), + new_state=CoordinatorState.STOPPING, + ), + # No state change; broadcast in any state so engines stay in sync. + Headers.SET_GENERATION_EPOCH: ControlTransition(allowed_from=_ALL_STATES, new_state=None), +} diff --git a/megatron/core/inference/disaggregation/mamba_reshard.py b/megatron/core/inference/disaggregation/mamba_reshard.py deleted file mode 100644 index 8a23735154a..00000000000 --- a/megatron/core/inference/disaggregation/mamba_reshard.py +++ /dev/null @@ -1,222 +0,0 @@ -# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. - -"""Heterogeneous TP/PP reshard of Mamba conv/ssm state between prefill and -decode shard layouts (the Mamba analog of the attention KV reshard).""" - -from __future__ import annotations - -from dataclasses import dataclass -from typing import List, Tuple - -from megatron.core.inference.disaggregation.utils import intersect - -# Channel bands of a Mamba layer's state, in the order the conv state -# concatenates them on its channel axis (x, B, C); ssm is the head axis. -# (name, lives_in_conv). conv bands share one tensor; ssm is its own tensor. -_CONV_BANDS = ("x", "B", "C") - - -@dataclass(frozen=True) -class MambaStateDims: - """The model's (global, unsharded) Mamba structural dims. - - These belong to the MambaMixer / model config -- carried as one unit (rather - than loose constants spread across the layout) so there's a single source - and they can't drift apart. The producer should read them straight from the - model config (e.g. ``ngroups = config.mamba_num_groups``) rather than - reverse-deriving from tensor shapes. TP shards ``nheads``/``ngroups``; the - rest are unsharded. - """ - - nheads: int - headdim: int - d_state: int - ngroups: int - d_conv: int - - -@dataclass(frozen=True) -class MambaShardLayout: - """One rank's Mamba-state ownership: which global layers + TP rank, plus the - model's structural dims (:class:`MambaStateDims`). Per-rank locals follow by - dividing by ``tp_size``.""" - - global_rank: int - tp_size: int - tp_rank: int - layer_start: int # global Mamba-layer index of this rank's first layer - num_layers: int # Mamba layers held locally (this PP stage) - dims: MambaStateDims - - def __post_init__(self) -> None: - # Wire reconstruction (MambaShardLayout(**dict)) hands ``dims`` as a - # plain dict; coerce it back to MambaStateDims. - if isinstance(self.dims, dict): - object.__setattr__(self, "dims", MambaStateDims(**self.dims)) - # TP shards heads and groups; both must divide evenly or the local - # conv/ssm band sizes truncate to the wrong (or zero) width silently. - if self.dims.nheads % self.tp_size != 0: - raise ValueError(f"nheads={self.dims.nheads} not divisible by tp_size={self.tp_size}") - if self.dims.ngroups % self.tp_size != 0: - raise ValueError(f"ngroups={self.dims.ngroups} not divisible by tp_size={self.tp_size}") - - # Convenience proxies onto the dims so callers read ``layout.headdim`` etc. - @property - def nheads(self) -> int: - """Global (unsharded) number of Mamba heads.""" - return self.dims.nheads - - @property - def headdim(self) -> int: - """Dimension of each Mamba head.""" - return self.dims.headdim - - @property - def d_state(self) -> int: - """SSM state size per head.""" - return self.dims.d_state - - @property - def ngroups(self) -> int: - """Global (unsharded) number of B/C groups.""" - return self.dims.ngroups - - @property - def d_conv(self) -> int: - """Convolution kernel width.""" - return self.dims.d_conv - - def mamba_shard_key(self) -> Tuple[int, int]: - """The Mamba shard this rank holds: ``(tp_rank, layer_start)``. Ranks - sharing a key hold identical state (e.g. EP/DP replicas of it).""" - return (self.tp_rank, self.layer_start) - - @property - def d_inner(self) -> int: - """Global inner dimension (nheads * headdim).""" - return self.dims.nheads * self.dims.headdim - - @property - def nheads_local(self) -> int: - """Number of Mamba heads held by this TP rank.""" - return self.dims.nheads // self.tp_size - - @property - def d_inner_local(self) -> int: - """Local inner dimension for this TP rank.""" - return self.d_inner // self.tp_size - - @property - def ngroups_local(self) -> int: - """Number of B/C groups held by this TP rank.""" - return self.dims.ngroups // self.tp_size - - @property - def conv_dim_local(self) -> int: - """Total local conv channel width (x + B + C bands).""" - return self.d_inner_local + 2 * self.ngroups_local * self.dims.d_state - - def layer_range(self) -> Tuple[int, int]: - """Global Mamba-layer range ``[lo, hi)`` owned by this rank.""" - return (self.layer_start, self.layer_start + self.num_layers) - - def _band(self, name: str) -> Tuple[int, int, int]: - """``(global_total, local_size, conv_local_offset)`` for a band. - - ``conv_local_offset`` is the band's start on the local conv channel - axis; for the ``ssm`` (head) band it is the start on the local head - axis (always 0, heads are the whole tensor).""" - if name == "x": - g = self.d_inner - return g, self.d_inner_local, 0 - if name == "B": - g = self.dims.ngroups * self.dims.d_state - return g, self.ngroups_local * self.dims.d_state, self.d_inner_local - if name == "C": - g = self.dims.ngroups * self.dims.d_state - return ( - g, - self.ngroups_local * self.dims.d_state, - self.d_inner_local + self.ngroups_local * self.dims.d_state, - ) - if name == "ssm": - return self.dims.nheads, self.nheads_local, 0 - raise KeyError(name) - - -@dataclass(frozen=True) -class MambaReshardTransfer: - """One sub-block move for the reshard. - - ``band`` is ``"x"``/``"B"``/``"C"`` (conv channel axis) or ``"ssm"`` (head - axis). ``src_layer``/``dst_layer`` are local layer indices on each side; - ``*_lo``/``*_hi`` are the local channel/head slice bounds. - """ - - src_rank: int - dst_rank: int - band: str - global_layer: int - src_layer: int - dst_layer: int - src_lo: int - src_hi: int - dst_lo: int - dst_hi: int - - @property - def is_conv(self) -> bool: - """True if this transfer targets the conv state; False for ssm.""" - return self.band in _CONV_BANDS - - -def plan_mamba_reshard( - src_layouts: List[MambaShardLayout], dst_layouts: List[MambaShardLayout] -) -> List[MambaReshardTransfer]: - """Plan the conv/ssm sub-block moves from the prefill (src) layouts to the - decode (dst) layouts. One transfer per (src rank, dst rank, global layer, - band) where both the layer ranges and the channel ranges overlap.""" - # Dedupe replica sources: ranks sharing (tp_rank, layer_start) hold identical - # Mamba state (e.g. EP/DP replicas), so source each shard from exactly one of - # them -- the smallest global_rank -- to avoid duplicate sends. - rep_rank: dict = {} - for s in src_layouts: - key = s.mamba_shard_key() - if key not in rep_rank or s.global_rank < rep_rank[key]: - rep_rank[key] = s.global_rank - source_ranks = set(rep_rank.values()) - - out: List[MambaReshardTransfer] = [] - for s in src_layouts: - if s.global_rank not in source_ranks: - continue - s_lr = s.layer_range() - for d in dst_layouts: - layer_ov = intersect(s_lr, d.layer_range()) - if layer_ov is None: - continue - for band in (*_CONV_BANDS, "ssm"): - _, s_size, s_off = s._band(band) - _, d_size, d_off = d._band(band) - s_glo = (s.tp_rank * s_size, s.tp_rank * s_size + s_size) - d_glo = (d.tp_rank * d_size, d.tp_rank * d_size + d_size) - chan_ov = intersect(s_glo, d_glo) - if chan_ov is None: - continue - lo, hi = chan_ov - for g in range(layer_ov[0], layer_ov[1]): - out.append( - MambaReshardTransfer( - src_rank=s.global_rank, - dst_rank=d.global_rank, - band=band, - global_layer=g, - src_layer=g - s.layer_start, - dst_layer=g - d.layer_start, - src_lo=s_off + (lo - s_glo[0]), - src_hi=s_off + (hi - s_glo[0]), - dst_lo=d_off + (lo - d_glo[0]), - dst_hi=d_off + (hi - d_glo[0]), - ) - ) - return out diff --git a/megatron/core/inference/disaggregation/ssm_reshard.py b/megatron/core/inference/disaggregation/ssm_reshard.py new file mode 100644 index 00000000000..9c06b9f70f9 --- /dev/null +++ b/megatron/core/inference/disaggregation/ssm_reshard.py @@ -0,0 +1,205 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Heterogeneous TP/PP reshard of SSM boundary-snapshot state between +prefill and decode shard layouts (the SSM analog of attention KV resharding). + +A snapshot's conv state packs three channel bands, [x | B | C], on one axis: +x is head-sharded (d_inner) and B/C are group-sharded (ngroups * d_state). +The recurrent state is head-sharded. ``plan_ssm_reshard`` emits one transfer per +(src rank, dst rank, global layer, band) whose layer and channel ranges +overlap; both sides compute the same plan from the same layout lists, so the +send and receive orders match. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import List, Tuple + +from megatron.core.inference.disaggregation.utils import intersect + +# Channel bands of an SSM layer's conv state, in the order the conv state +# concatenates them on its channel axis. The recurrent band is the head axis +# of its own tensor. +_CONV_BANDS = ("x", "B", "C") + + +@dataclass(frozen=True) +class SSMStateDims: + """The model's global (unsharded) SSM structural dimensions. + + Carried as one unit so there is a single source and the dims cannot drift + apart. The producer should read them from the model config rather than + deriving them from tensor shapes. TP shards nheads/ngroups; the rest are + unsharded. + """ + + nheads: int + headdim: int + d_state: int + ngroups: int + d_conv: int + + +@dataclass(frozen=True) +class SSMShardLayout: + """One rank's SSM-state ownership: which global layers and TP rank, + plus the model's structural dims. Per-rank local sizes follow by dividing + by tp_size.""" + + global_rank: int + tp_size: int + tp_rank: int + layer_start: int # global SSM-layer index of this rank's first layer + num_layers: int # SSM layers held locally (this PP stage) + dims: SSMStateDims + + def __post_init__(self) -> None: + # Wire reconstruction (SSMShardLayout(**dict)) hands dims as a plain + # dict; coerce it back to SSMStateDims. + if isinstance(self.dims, dict): + object.__setattr__(self, "dims", SSMStateDims(**self.dims)) + # TP shards heads and groups; both must divide evenly or the local + # band widths are wrong. + if self.dims.nheads % self.tp_size != 0: + raise ValueError(f"nheads={self.dims.nheads} not divisible by tp_size={self.tp_size}") + if self.dims.ngroups % self.tp_size != 0: + raise ValueError(f"ngroups={self.dims.ngroups} not divisible by tp_size={self.tp_size}") + + @property + def d_inner(self) -> int: + """Global inner dimension (nheads * headdim).""" + return self.dims.nheads * self.dims.headdim + + @property + def nheads_local(self) -> int: + """SSM heads held by this TP rank.""" + return self.dims.nheads // self.tp_size + + @property + def d_inner_local(self) -> int: + """Local inner dimension for this TP rank.""" + return self.d_inner // self.tp_size + + @property + def ngroups_local(self) -> int: + """B/C groups held by this TP rank.""" + return self.dims.ngroups // self.tp_size + + @property + def conv_dim_local(self) -> int: + """Total local conv channel width (x + B + C bands).""" + return self.d_inner_local + 2 * self.ngroups_local * self.dims.d_state + + def shard_key(self) -> Tuple[int, int]: + """The SSM shard this rank holds: (tp_rank, layer_start). Ranks + sharing a key hold identical state (e.g. EP/DP replicas).""" + return (self.tp_rank, self.layer_start) + + def layer_range(self) -> Tuple[int, int]: + """Global SSM-layer range [lo, hi) owned by this rank.""" + return (self.layer_start, self.layer_start + self.num_layers) + + def band(self, name: str) -> Tuple[int, int, int]: + """Return (global_total, local_size, local_offset) for a band. + + local_offset is the band's start on the local conv channel axis; for + the "recurrent" (head) band it is the start on the local head axis + (always 0, heads are the whole tensor). + """ + if name == "x": + return self.d_inner, self.d_inner_local, 0 + if name == "B": + g = self.dims.ngroups * self.dims.d_state + return g, self.ngroups_local * self.dims.d_state, self.d_inner_local + if name == "C": + g = self.dims.ngroups * self.dims.d_state + return ( + g, + self.ngroups_local * self.dims.d_state, + self.d_inner_local + self.ngroups_local * self.dims.d_state, + ) + if name == "recurrent": + return self.dims.nheads, self.nheads_local, 0 + raise KeyError(name) + + +@dataclass(frozen=True) +class SSMReshardTransfer: + """One sub-block move of the snapshot reshard. + + band is "x"/"B"/"C" (conv channel axis) or "recurrent" (head axis). + src_layer/dst_layer are local layer indices on each side; *_lo/*_hi are + the local channel or head slice bounds. + """ + + src_rank: int + dst_rank: int + band: str + global_layer: int + src_layer: int + dst_layer: int + src_lo: int + src_hi: int + dst_lo: int + dst_hi: int + + @property + def is_conv(self) -> bool: + """True if this transfer targets the conv state; False for recurrent.""" + return self.band in _CONV_BANDS + + +def plan_ssm_reshard( + src_layouts: List[SSMShardLayout], dst_layouts: List[SSMShardLayout] +) -> List[SSMReshardTransfer]: + """Plan the conv/recurrent sub-block moves from the prefill (src) layouts to + the decode (dst) layouts: one transfer per (src rank, dst rank, global + layer, band) where both the layer ranges and the channel ranges overlap. + + Ranks sharing (tp_rank, layer_start) hold identical SSM state (e.g. + EP/DP replicas), so each shard is sourced from exactly one of them, the + smallest global_rank. Deterministic given the layout lists, so both sides + enumerate the same transfers in the same order. + """ + rep_rank: dict = {} + for s in src_layouts: + key = s.shard_key() + if key not in rep_rank or s.global_rank < rep_rank[key]: + rep_rank[key] = s.global_rank + source_ranks = set(rep_rank.values()) + + out: List[SSMReshardTransfer] = [] + for s in src_layouts: + if s.global_rank not in source_ranks: + continue + s_lr = s.layer_range() + for d in dst_layouts: + layer_ov = intersect(s_lr, d.layer_range()) + if layer_ov is None: + continue + for band in (*_CONV_BANDS, "recurrent"): + _, s_size, s_off = s.band(band) + _, d_size, d_off = d.band(band) + s_glo = (s.tp_rank * s_size, s.tp_rank * s_size + s_size) + d_glo = (d.tp_rank * d_size, d.tp_rank * d_size + d_size) + chan_ov = intersect(s_glo, d_glo) + if chan_ov is None: + continue + lo, hi = chan_ov + for g in range(layer_ov[0], layer_ov[1]): + out.append( + SSMReshardTransfer( + src_rank=s.global_rank, + dst_rank=d.global_rank, + band=band, + global_layer=g, + src_layer=g - s.layer_start, + dst_layer=g - d.layer_start, + src_lo=s_off + (lo - s_glo[0]), + src_hi=s_off + (hi - s_glo[0]), + dst_lo=d_off + (lo - d_glo[0]), + dst_hi=d_off + (hi - d_glo[0]), + ) + ) + return out diff --git a/megatron/core/inference/disaggregation/transfer_backends/__init__.py b/megatron/core/inference/disaggregation/transfer_backends/__init__.py new file mode 100644 index 00000000000..262675e55c0 --- /dev/null +++ b/megatron/core/inference/disaggregation/transfer_backends/__init__.py @@ -0,0 +1,7 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""KV transfer backends for disaggregated inference.""" + +from .base import KVTransportBackend, construct_kv_transfer_backend_class + +__all__ = ["KVTransportBackend", "construct_kv_transfer_backend_class"] diff --git a/megatron/core/inference/disaggregation/transfer_backends/base.py b/megatron/core/inference/disaggregation/transfer_backends/base.py new file mode 100644 index 00000000000..fa8df492e44 --- /dev/null +++ b/megatron/core/inference/disaggregation/transfer_backends/base.py @@ -0,0 +1,212 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""KV transfer backend registry and the buffer geometry shared by backends. + +Backends are selected explicitly by the caller's launcher configuration, +never from the environment. +""" + +from __future__ import annotations + +from dataclasses import asdict, dataclass +from typing import Any, Optional + +import torch + +from megatron.core.inference.disaggregation.kv_reshard import KVShardLayout +from megatron.core.inference.disaggregation.ssm_reshard import SSMShardLayout + +KVTransportBackend = Any + + +def construct_kv_transfer_backend_class(name: str) -> KVTransportBackend: + """Return the backend class registered under ``name``.""" + + normalized = name.lower().replace("_", "-") + if normalized == "nixl": + from .nixl import NixlTransferBackend + + return NixlTransferBackend + if normalized == "nccl": + from .nccl import NcclTransferBackend + + return NcclTransferBackend + raise ValueError("Unsupported KV transfer backend %r; expected 'nixl' or 'nccl'." % name) + + +@dataclass +class BufferGeometry: + """Address geometry of one registered paged buffer. + + Each (outer, block) pair is one contiguous slice; the outer stride skips + over the full block pool for that outer index. ``layout`` is the canonical + KV shard layout when the buffer is a KV cache; SSM pools carry their + typed layout separately. + """ + + buf_ptr: int + element_size: int + device_id: int + blocks_axis: int + num_blocks: int + num_outer: int + bytes_per_slice: int + outer_stride_bytes: int + heads_per_partition: Optional[int] + head_dim: Optional[int] + tokens_per_block: Optional[int] + layout: Optional[KVShardLayout] + + +def compute_buffer_geometry( + memory_buffer: torch.Tensor, + expected_num_blocks: int, + *, + backend_name: str, + tp_size: Optional[int] = None, + tp_rank: Optional[int] = None, + num_kv_heads_global: Optional[int] = None, + heads_per_partition: Optional[int] = None, + head_dim: Optional[int] = None, + tokens_per_block: Optional[int] = None, + global_rank: Optional[int] = None, + pp_size: Optional[int] = None, + pp_rank: Optional[int] = None, + num_layers_global: Optional[int] = None, + layer_start: Optional[int] = None, + layer_end: Optional[int] = None, + ssm_layout: Optional[SSMShardLayout] = None, + ssm_state_kind: Optional[str] = None, +) -> BufferGeometry: + """Locate the blocks axis, derive the slice strides, and validate the + canonical KV layout when the full geometry is provided. + + Shared by every transfer backend so they agree on addressing and on the + exported metadata schema. The inference KV layout is [2, L, B, T, H, d]. + """ + if (ssm_layout is None) != (ssm_state_kind is None): + raise ValueError("ssm_layout and ssm_state_kind must be provided together") + if ssm_state_kind not in (None, "conv", "recurrent"): + raise ValueError("ssm_state_kind must be 'conv' or 'recurrent'") + + layout_capable = ( + None + not in ( + global_rank, + tp_size, + tp_rank, + pp_size, + pp_rank, + num_layers_global, + num_kv_heads_global, + heads_per_partition, + head_dim, + tokens_per_block, + layer_start, + layer_end, + ) + and heads_per_partition * tp_size == num_kv_heads_global + ) + + shape = list(memory_buffer.shape) + candidates = [i for i, dim in enumerate(shape) if dim == expected_num_blocks] + if not candidates: + raise RuntimeError( + f"{backend_name}: no axis in memory_buffer shape {shape} matches " + f"expected_num_blocks={expected_num_blocks}. Layout is unrecognized; " + "bug in caller or new Megatron tensor shape." + ) + if len(candidates) > 1: + raise RuntimeError( + f"{backend_name}: ambiguous blocks axis in shape {shape} " + f"(expected_num_blocks={expected_num_blocks} matches multiple axes " + f"{candidates}). Caller must pass a more distinctive value." + ) + blocks_axis = candidates[0] + + elements_per_slice = 1 + for dim in shape[blocks_axis + 1 :]: + elements_per_slice *= dim + element_size = memory_buffer.element_size() + bytes_per_slice = element_size * elements_per_slice + num_outer = 1 + for dim in shape[:blocks_axis]: + num_outer *= dim + + layout = None + if layout_capable: + layout = KVShardLayout( + num_layers=int(num_layers_global), + num_heads=int(num_kv_heads_global), + tp_size=int(tp_size), + tp_rank=int(tp_rank), + pp_size=int(pp_size), + pp_rank=int(pp_rank), + global_rank=int(global_rank), + layer_start=int(layer_start), + num_local_layers=int(layer_end) - int(layer_start), + ) + if blocks_axis != 2: + raise ValueError("inference KV transfers require the [2, L, B, T, H, d] layout") + if layout.local_num_heads() != heads_per_partition: + raise ValueError( + "heads_per_partition does not match the canonical KV layout: " + f"{heads_per_partition} vs {layout.local_num_heads()}" + ) + if num_outer % layout.local_num_layers() != 0: + raise ValueError( + f"num_outer={num_outer} is not divisible by local layers=" + f"{layout.local_num_layers()}" + ) + if layout is not None and ssm_layout is not None: + raise ValueError("a transfer backend cannot have both KV and SSM layouts") + + return BufferGeometry( + buf_ptr=memory_buffer.data_ptr(), + element_size=element_size, + device_id=memory_buffer.device.index if memory_buffer.is_cuda else 0, + blocks_axis=blocks_axis, + num_blocks=expected_num_blocks, + num_outer=num_outer, + bytes_per_slice=bytes_per_slice, + outer_stride_bytes=expected_num_blocks * bytes_per_slice, + heads_per_partition=heads_per_partition, + head_dim=head_dim, + tokens_per_block=tokens_per_block, + layout=layout, + ) + + +def export_geometry_meta(geometry: BufferGeometry, ssm_layout=None) -> dict: + """The wire schema shared by every backend's export_meta.""" + meta = { + "base_addr": geometry.buf_ptr, + "outer_stride_bytes": geometry.outer_stride_bytes, + "device_id": geometry.device_id, + "num_outer": geometry.num_outer, + "bytes_per_slice": geometry.bytes_per_slice, + "blocks_axis": geometry.blocks_axis, + "num_blocks": geometry.num_blocks, + "heads_per_partition": geometry.heads_per_partition, + "head_dim": geometry.head_dim, + "tokens_per_block": geometry.tokens_per_block, + "element_size": geometry.element_size, + } + if geometry.layout is not None: + layer_start, layer_end = geometry.layout.layer_range() + meta.update( + { + "global_rank": geometry.layout.global_rank, + "tp_size": geometry.layout.tp_size, + "tp_rank": geometry.layout.tp_rank, + "pp_size": geometry.layout.pp_size, + "pp_rank": geometry.layout.pp_rank, + "num_layers_global": geometry.layout.num_layers, + "num_kv_heads_global": geometry.layout.num_heads, + "layer_start": layer_start, + "layer_end": layer_end, + } + ) + if ssm_layout is not None: + meta["ssm_layout"] = asdict(ssm_layout) + return meta diff --git a/megatron/core/inference/disaggregation/transfer_backends/nccl.py b/megatron/core/inference/disaggregation/transfer_backends/nccl.py new file mode 100644 index 00000000000..96350fb54bf --- /dev/null +++ b/megatron/core/inference/disaggregation/transfer_backends/nccl.py @@ -0,0 +1,293 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Two-sided (NCCL) KV transfer backend for disaggregated prefill/decode. + +Unlike the one-sided NIXL backend, both peers participate: the decode posts +receives when the hand-off request arrives (begin_pull_blocks) and the prefill +posts the matching sends when the coordinator's SEND_KV names the decode +instance (begin_push_blocks). Both sides enumerate the same reshard plan in +the same deterministic order, so the point-to-point operations match by post +order per peer pair. Data moves straight out of the prefill's pinned blocks; +there is no staging copy on the send side. +""" + +from __future__ import annotations + +import logging +from typing import Any, Dict, List, Optional + +import torch +import torch.distributed as dist + +from megatron.core.inference.disaggregation.kv_reshard import KVShardLayout, plan_kv_reshard +from megatron.core.inference.disaggregation.ssm_reshard import SSMShardLayout, plan_ssm_reshard +from megatron.core.inference.disaggregation.transfer_backends.base import ( + compute_buffer_geometry, + export_geometry_meta, +) +from megatron.core.inference.disaggregation.utils import transfer_peer_records + +logger = logging.getLogger(__name__) + + +class NcclTransferHandle: + """Pollable handle for one batched NCCL transfer. + + Receives land in temporary contiguous buffers; on completion the handle + runs its scatter closures once to place the data into the paged buffers. + Send handles keep the gathered source slices alive until reaped. + """ + + def __init__(self, works: List[Any], keepalive: List[torch.Tensor], scatters: List[Any]): + self._works = works + self._keepalive = keepalive + self._scatters = scatters + self._done = not works and not scatters + + def poll(self) -> bool: + """Return True if the transfer has settled, scattering received data + into the paged buffers on first completion.""" + if self._done: + return True + if not all(w.is_completed() for w in self._works): + return False + self._finish() + return True + + def wait(self) -> None: + """Block until the transfer completes, then scatter.""" + for w in self._works: + w.wait() + self._finish() + + def _finish(self) -> None: + if self._done: + return + with torch.inference_mode(): + for scatter in self._scatters: + scatter() + self._keepalive.clear() + self._works = [] + self._done = True + + +def _make_copy(view: torch.Tensor, buf: torch.Tensor): + def _copy(): + view.copy_(buf.view(view.shape)) + + return _copy + + +def _kv_layout_from_meta(meta: Dict[str, Any]) -> KVShardLayout: + """Rebuild a peer's KVShardLayout from its exported metadata.""" + return KVShardLayout( + num_layers=int(meta["num_layers_global"]), + num_heads=int(meta["num_kv_heads_global"]), + tp_size=int(meta["tp_size"]), + tp_rank=int(meta["tp_rank"]), + pp_size=int(meta["pp_size"]), + pp_rank=int(meta["pp_rank"]), + global_rank=int(meta["global_rank"]), + layer_start=int(meta["layer_start"]), + num_local_layers=int(meta["layer_end"]) - int(meta["layer_start"]), + ) + + +class NcclTransferBackend: + """Per-buffer NCCL transport over the default process group. + + Mirrors the NIXL backend's construction and metadata schema so the + hand-off layer treats the two interchangeably; only the transfer calls + differ (two-sided matched send/recv instead of one-sided reads). + """ + + name = "nccl" + is_push = True + + def __init__( + self, + agent_name: str, + memory_buffer: torch.Tensor, + expected_num_blocks: int, + tp_size: Optional[int] = None, + tp_rank: Optional[int] = None, + num_kv_heads_global: Optional[int] = None, + heads_per_partition: Optional[int] = None, + head_dim: Optional[int] = None, + tokens_per_block: Optional[int] = None, + global_rank: Optional[int] = None, + pp_size: Optional[int] = None, + pp_rank: Optional[int] = None, + num_layers_global: Optional[int] = None, + layer_start: Optional[int] = None, + layer_end: Optional[int] = None, + ssm_layout: Optional[SSMShardLayout] = None, + ssm_state_kind: Optional[str] = None, + ): + if not (dist.is_available() and dist.is_initialized()): + raise RuntimeError( + "NcclTransferBackend requires torch.distributed to be initialized; " + "the prefill and decode workers must share a process group." + ) + self.agent_name = agent_name + self._memory_buffer = memory_buffer + geometry = compute_buffer_geometry( + memory_buffer, + expected_num_blocks, + backend_name="NcclTransferBackend", + tp_size=tp_size, + tp_rank=tp_rank, + num_kv_heads_global=num_kv_heads_global, + heads_per_partition=heads_per_partition, + head_dim=head_dim, + tokens_per_block=tokens_per_block, + global_rank=global_rank, + pp_size=pp_size, + pp_rank=pp_rank, + num_layers_global=num_layers_global, + layer_start=layer_start, + layer_end=layer_end, + ssm_layout=ssm_layout, + ssm_state_kind=ssm_state_kind, + ) + self._geometry = geometry + self._layout = geometry.layout + self._ssm_layout = ssm_layout + self._ssm_state_kind = ssm_state_kind + logger.info( + "NcclTransferBackend[%s] over %d-block buffer (rank=%d, shape=%s)", + agent_name, + geometry.num_blocks, + dist.get_rank(), + list(memory_buffer.shape), + ) + + def export_meta(self) -> Dict[str, Any]: + """The shared geometry schema plus this rank's NCCL address.""" + meta = export_geometry_meta(self._geometry, self._ssm_layout) + meta["transport"] = "nccl" + meta["nccl_rank"] = dist.get_rank() + return meta + + # --- shared enumeration ------------------------------------------------- + def _kv_transfers(self, peer_records, mine_is_src: bool): + """Yield (peer_meta, layers, heads) for this rank's part of the KV + reshard plan, in deterministic plan order.""" + sources: list = [] + peers_by_rank: dict = {} + for meta, blocks in peer_records: + layout = _kv_layout_from_meta(meta) + if layout.global_rank in peers_by_rank: + raise ValueError(f"duplicate peer global_rank={layout.global_rank} in KV metadata") + peers_by_rank[layout.global_rank] = meta + sources.append(layout) + if mine_is_src: + plan = plan_kv_reshard([self._layout], sources) + else: + plan = plan_kv_reshard(sources, [self._layout]) + for transfer in plan: + peer_rank = transfer.dst_rank if mine_is_src else transfer.src_rank + meta = peers_by_rank[peer_rank] + if mine_is_src: + layers = transfer.src_layer_slice(self._layout) + heads = transfer.src_head_slice(self._layout) + else: + layers = transfer.dst_layer_slice(self._layout) + heads = transfer.dst_head_slice(self._layout) + yield meta, layers, heads + + def _ssm_transfers(self, peer_records, mine_is_src: bool): + """Yield (peer_meta, lo, hi) band slices of this rank's SSM state, + in deterministic plan order.""" + for meta, _ in peer_records: + raw_layout = meta.get("ssm_layout") + if not isinstance(raw_layout, dict): + raise ValueError("peer metadata is missing ssm_layout") + peer_layout = SSMShardLayout(**raw_layout) + if mine_is_src: + plan = plan_ssm_reshard([self._ssm_layout], [peer_layout]) + else: + plan = plan_ssm_reshard([peer_layout], [self._ssm_layout]) + for t in plan: + if t.is_conv != (self._ssm_state_kind == "conv"): + continue + if mine_is_src: + yield meta, t.src_layer, t.src_lo, t.src_hi + else: + yield meta, t.dst_layer, t.dst_lo, t.dst_hi + + def _kv_block_view(self, block_id: int, layers: slice, heads: slice) -> torch.Tensor: + """One block's (kv, layer, token, head, dim) fragment in the + [2, L, B, T, H, d] paged buffer.""" + return self._memory_buffer[:, layers, block_id, :, heads, :] + + # --- decode side --------------------------------------------------------- + def begin_pull_blocks( + self, peer_meta: Any, src_block_ids: List[int], dst_block_ids: List[int] + ) -> NcclTransferHandle: + """Post the receives matching the prefill's sends; the handle scatters + into the destination blocks (or SSM slots) on completion.""" + if not dst_block_ids: + return NcclTransferHandle([], [], []) + records = transfer_peer_records(peer_meta, src_block_ids) + + ops: List[Any] = [] + buffers: List[torch.Tensor] = [] + scatters: List[Any] = [] + device = self._memory_buffer.device + dtype = self._memory_buffer.dtype + + if self._ssm_layout is not None: + for meta, layer, lo, hi in self._ssm_transfers(records, mine_is_src=False): + for slot in dst_block_ids: + view = self._memory_buffer[layer, int(slot), lo:hi] + buf = torch.empty(view.shape, dtype=dtype, device=device) + buffers.append(buf) + ops.append(dist.P2POp(dist.irecv, buf, int(meta["nccl_rank"]))) + scatters.append(_make_copy(view, buf)) + else: + geo = self._geometry + for meta, layers, heads in self._kv_transfers(records, mine_is_src=False): + n_layers = layers.stop - layers.start + n_heads = heads.stop - heads.start + for block in dst_block_ids: + buf = torch.empty( + (2, n_layers, geo.tokens_per_block, n_heads, geo.head_dim), + dtype=dtype, + device=device, + ) + buffers.append(buf) + ops.append(dist.P2POp(dist.irecv, buf, int(meta["nccl_rank"]))) + scatters.append(_make_copy(self._kv_block_view(int(block), layers, heads), buf)) + + works = dist.batch_isend_irecv(ops) if ops else [] + return NcclTransferHandle(works, buffers, scatters) + + # --- prefill side ---------------------------------------------------------- + def begin_push_blocks(self, peer_meta: Any, src_block_ids: List[int]) -> NcclTransferHandle: + """Post the sends matching the decode's receives, straight out of the + pinned source blocks (or SSM slots). `peer_meta` is the decode + instance's per-rank metadata in the same nested shape as a hand-off's + kv_meta.""" + if not src_block_ids: + return NcclTransferHandle([], [], []) + records = transfer_peer_records(peer_meta, []) + + ops: List[Any] = [] + keep: List[torch.Tensor] = [] + + if self._ssm_layout is not None: + for meta, layer, lo, hi in self._ssm_transfers(records, mine_is_src=True): + for slot in src_block_ids: + sub = self._memory_buffer[layer, int(slot), lo:hi].contiguous() + keep.append(sub) + ops.append(dist.P2POp(dist.isend, sub, int(meta["nccl_rank"]))) + else: + for meta, layers, heads in self._kv_transfers(records, mine_is_src=True): + for block in src_block_ids: + sub = self._kv_block_view(int(block), layers, heads).contiguous() + keep.append(sub) + ops.append(dist.P2POp(dist.isend, sub, int(meta["nccl_rank"]))) + + works = dist.batch_isend_irecv(ops) if ops else [] + return NcclTransferHandle(works, keep, []) diff --git a/megatron/core/inference/disaggregation/transfer_backends/nixl.py b/megatron/core/inference/disaggregation/transfer_backends/nixl.py new file mode 100644 index 00000000000..af355e25045 --- /dev/null +++ b/megatron/core/inference/disaggregation/transfer_backends/nixl.py @@ -0,0 +1,625 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Direct NIXL backend for disaggregated prefill/decode KV transfer. + +Each rank registers its paged KV buffer once, exports NIXL peer metadata, and +the decode side pulls source block ranges directly into its local KV blocks. + +Backend selection belongs in ``transfer_backends.base`` and is supplied +explicitly by the launcher. +""" + +from __future__ import annotations + +import base64 +import logging +import os +import time +from dataclasses import asdict, dataclass +from typing import Any, Dict, List, Optional + +import torch + +from megatron.core.inference.disaggregation.kv_reshard import KVShardLayout, plan_kv_reshard +from megatron.core.inference.disaggregation.ssm_reshard import SSMShardLayout, plan_ssm_reshard +from megatron.core.inference.disaggregation.transfer_backends.base import compute_buffer_geometry +from megatron.core.inference.disaggregation.utils import transfer_peer_records + +logger = logging.getLogger(__name__) + +try: + from nixl._api import nixl_agent # type: ignore[import-not-found] + + _HAVE_NIXL = True +except ImportError: + nixl_agent = None # type: ignore[assignment] + _HAVE_NIXL = False + + +# NIXL exposes polling, not a blocking wait. A long stall usually means peer or +# fabric failure, so cap the wait. +_POLL_INTERVAL_S = 0.0005 # 0.5 ms +_POLL_TIMEOUT_S = 30.0 + + +@dataclass +class NixlPullHandle: + """Pollable handle for one logical pull made of one or more NIXL transfers.""" + + agent: Any + xfers: List[Any] + contexts: List[str] + submitted_at: float + timeout_s: float = _POLL_TIMEOUT_S + done: bool = False + error: Optional[str] = None + + def poll(self) -> bool: + """Return True if every transfer has settled, without blocking.""" + if self.done: + if self.error is not None: + raise RuntimeError(self.error) + return True + if not self.xfers: + self.done = True + return True + + errors: List[str] = [] + pending: List[str] = [] + for xfer, ctx in zip(self.xfers, self.contexts): + state = self.agent.check_xfer_state(xfer) + if state == "DONE": + continue + if state == "ERR": + errors.append(ctx) + continue + pending.append(f"{ctx}: {state}") + + if not pending: + self.done = True + if errors: + self.error = f"NIXL transfer failed ({', '.join(errors)})" + raise RuntimeError(self.error) + return True + if time.perf_counter() - self.submitted_at > self.timeout_s: + raise TimeoutError( + f"NIXL transfer timed out after {self.timeout_s}s; pending={pending}" + ) + return False + + def wait(self) -> None: + """Block until the transfer completes; NIXL has no blocking wait, so + poll with a short sleep to avoid monopolizing a CPU core.""" + while not self.poll(): + time.sleep(_POLL_INTERVAL_S) + + +class NixlTransferBackend: + """Per-rank NIXL agent owning a registration over the paged KV buffer. + + Per-block transfers are descriptor ranges over that registration. Peer + metadata is exchanged by the control plane and registered lazily on first + pull. + """ + + name = "nixl" + + def __init__( + self, + agent_name: str, + memory_buffer: torch.Tensor, + expected_num_blocks: int, + tp_size: Optional[int] = None, + tp_rank: Optional[int] = None, + num_kv_heads_global: Optional[int] = None, + heads_per_partition: Optional[int] = None, + head_dim: Optional[int] = None, + tokens_per_block: Optional[int] = None, + global_rank: Optional[int] = None, + pp_size: Optional[int] = None, + pp_rank: Optional[int] = None, + num_layers_global: Optional[int] = None, + layer_start: Optional[int] = None, + layer_end: Optional[int] = None, + ssm_layout: Optional[SSMShardLayout] = None, + ssm_state_kind: Optional[str] = None, + ): + if not _HAVE_NIXL: + raise RuntimeError( + "NixlTransferBackend requires the nixl Python package. Install the " + "NIXL runtime and `pip install nixl` before launching " + "disaggregated workers." + ) + self.agent_name = agent_name + self._memory_buffer = memory_buffer + + # Addressing geometry shared with the other backends. + geometry = compute_buffer_geometry( + memory_buffer, + expected_num_blocks, + backend_name="NixlTransferBackend", + tp_size=tp_size, + tp_rank=tp_rank, + num_kv_heads_global=num_kv_heads_global, + heads_per_partition=heads_per_partition, + head_dim=head_dim, + tokens_per_block=tokens_per_block, + global_rank=global_rank, + pp_size=pp_size, + pp_rank=pp_rank, + num_layers_global=num_layers_global, + layer_start=layer_start, + layer_end=layer_end, + ssm_layout=ssm_layout, + ssm_state_kind=ssm_state_kind, + ) + self._geometry = geometry + shape = list(memory_buffer.shape) + self._buf_ptr = geometry.buf_ptr + self._element_size = geometry.element_size + self._device_id = geometry.device_id + self._outer_stride_bytes = geometry.outer_stride_bytes + self._num_outer = geometry.num_outer + self._bytes_per_slice = geometry.bytes_per_slice + self._blocks_axis = geometry.blocks_axis + self._num_blocks = geometry.num_blocks + self._heads_per_partition = geometry.heads_per_partition + self._head_dim = geometry.head_dim + self._tokens_per_block = geometry.tokens_per_block + self._layout = geometry.layout + self._ssm_layout = ssm_layout + self._ssm_state_kind = ssm_state_kind + + # Configure UCX before agent construction. Avoid TCP for VRAM addresses; + # operators may override this by setting UCX_TLS before launch. + os.environ.setdefault("UCX_TLS", "cuda_ipc,cuda_copy,cma,shm,self") + # Explicit registration makes the UCX memtype cache unnecessary and + # avoids stale VRAM/host classifications. + os.environ.setdefault("UCX_MEMTYPE_CACHE", "n") + + self._agent = nixl_agent(agent_name) + self._reg_handle = self._agent.register_memory(memory_buffer) + + # Base64 keeps NIXL metadata safe for msgpack/json control messages. + self._agent_metadata = self._agent.get_agent_metadata() + + # Peer agent_name -> id returned by add_remote_agent. + self._known_peers: Dict[str, Any] = {} + + logger.info( + "NixlTransferBackend[%s] registered %d-block buffer " + "(blocks_axis=%d, %d outer-slices/block × %d bytes/slice = " + "%d bytes/block, device=%d, shape=%s)", + agent_name, + self._num_blocks, + self._blocks_axis, + self._num_outer, + self._bytes_per_slice, + self._num_outer * self._bytes_per_slice, + self._device_id, + shape, + ) + + def export_meta(self) -> Dict[str, Any]: + """Return JSON/msgpack-safe metadata for shipping to a decode peer. + + Layout fields describe the scatter-gather address ranges needed to pull + source blocks into decode-owned blocks. + """ + meta = { + "agent_name": self.agent_name, + "agent_metadata_b64": base64.b64encode(self._agent_metadata).decode("ascii"), + "base_addr": self._buf_ptr, + "outer_stride_bytes": self._outer_stride_bytes, + "device_id": self._device_id, + "num_outer": self._num_outer, + "bytes_per_slice": self._bytes_per_slice, + "blocks_axis": self._blocks_axis, + "num_blocks": self._num_blocks, + "heads_per_partition": self._heads_per_partition, + "head_dim": self._head_dim, + "tokens_per_block": self._tokens_per_block, + "element_size": self._element_size, + } + if self._layout is not None: + layer_start, layer_end = self._layout.layer_range() + meta.update( + { + "global_rank": self._layout.global_rank, + "tp_size": self._layout.tp_size, + "tp_rank": self._layout.tp_rank, + "pp_size": self._layout.pp_size, + "pp_rank": self._layout.pp_rank, + "num_layers_global": self._layout.num_layers, + "num_kv_heads_global": self._layout.num_heads, + "layer_start": layer_start, + "layer_end": layer_end, + } + ) + if self._ssm_layout is not None: + meta["ssm_layout"] = asdict(self._ssm_layout) + return meta + + def _ensure_peer_registered(self, peer_meta: Dict[str, Any]) -> str: + """Register the peer with NIXL on first use; return its agent id.""" + peer_name = peer_meta["agent_name"] + existing = self._known_peers.get(peer_name) + if existing is not None: + return existing + metadata_b64 = peer_meta.get("agent_metadata_b64") + if not metadata_b64: + raise ValueError(f"peer_meta for {peer_name!r} is missing agent_metadata_b64") + peer_id = self._agent.add_remote_agent(base64.b64decode(metadata_b64)) + resolved = peer_id if peer_id else peer_name + self._known_peers[peer_name] = resolved + logger.info("NixlTransferBackend[%s] registered peer %s", self.agent_name, peer_name) + return resolved + + def _validate_peer( + self, + meta: Dict[str, Any], + src_block_ids: List[int], + dst_block_ids: List[int], + *, + matched_layout: bool = False, + ) -> None: + """Validate block mappings and physical transfer compatibility.""" + + if len(src_block_ids) != len(dst_block_ids): + raise ValueError( + f"source/destination block_id length mismatch for peer " + f"{meta.get('agent_name')!r}: {len(src_block_ids)} vs {len(dst_block_ids)}" + ) + for block in src_block_ids: + if not 0 <= block < int(meta["num_blocks"]): + raise ValueError(f"source block {block} is outside pool [0, {meta['num_blocks']})") + for block in dst_block_ids: + if not 0 <= block < self._num_blocks: + raise ValueError( + f"destination block {block} is outside pool [0, {self._num_blocks})" + ) + + local = { + "head_dim": self._head_dim, + "tokens_per_block": self._tokens_per_block, + "element_size": self._element_size, + "num_outer": self._num_outer, + "bytes_per_slice": self._bytes_per_slice, + "blocks_axis": self._blocks_axis, + "heads_per_partition": self._heads_per_partition, + } + fields = ["head_dim", "tokens_per_block", "element_size"] + if matched_layout: + fields.extend(["num_outer", "bytes_per_slice", "blocks_axis", "heads_per_partition"]) + mismatches = [ + f"{field}: peer={meta.get(field)} local={local[field]}" + for field in fields + if meta.get(field) is not None + and local[field] is not None + and meta.get(field) != local[field] + ] + if mismatches: + kind = "matched-layout" if matched_layout else "transfer" + raise ValueError(f"{kind} geometry mismatch: {', '.join(mismatches)}") + + @staticmethod + def _kv_layout_from_meta(meta: Dict[str, Any]) -> KVShardLayout: + """Reconstruct a main-planner KV layout from peer wire metadata.""" + + keys = ( + "global_rank", + "tp_size", + "tp_rank", + "pp_size", + "pp_rank", + "num_layers_global", + "num_kv_heads_global", + "layer_start", + "layer_end", + ) + missing = [key for key in keys if meta.get(key) is None] + if missing: + raise ValueError(f"peer metadata missing KV layout fields: {missing}") + return KVShardLayout( + num_layers=int(meta["num_layers_global"]), + num_heads=int(meta["num_kv_heads_global"]), + tp_size=int(meta["tp_size"]), + tp_rank=int(meta["tp_rank"]), + pp_size=int(meta["pp_size"]), + pp_rank=int(meta["pp_rank"]), + global_rank=int(meta["global_rank"]), + layer_start=int(meta["layer_start"]), + num_local_layers=int(meta["layer_end"]) - int(meta["layer_start"]), + ) + + def begin_pull_blocks( + self, peer_meta: Any, src_block_ids: List[int], dst_block_ids: List[int] + ) -> NixlPullHandle: + """Submit a pull and return a handle that can be polled later.""" + if not isinstance(peer_meta, dict) or "pp_metas" not in peer_meta: + if not src_block_ids and not dst_block_ids: + return NixlPullHandle( + agent=self._agent, + xfers=[], + contexts=[], + submitted_at=time.perf_counter(), + done=True, + ) + + xfers: List[Any] = [] + contexts: List[str] = [] + submitted_at = time.perf_counter() + try: + if self._ssm_layout is not None: + state_kind = self._ssm_state_kind + assert state_kind is not None + width = ( + self._ssm_layout.conv_dim_local + if state_kind == "conv" + else self._ssm_layout.nheads_local + ) + if ( + self._heads_per_partition != width + or self._num_outer != self._ssm_layout.num_layers + or self._blocks_axis != 1 + ): + raise ValueError(f"local {state_kind} geometry does not match its SSM layout") + + sources = [] + peers_by_rank = {} + for meta, blocks in transfer_peer_records(peer_meta, src_block_ids): + raw_layout = meta.get("ssm_layout") + if not isinstance(raw_layout, dict): + raise ValueError("peer metadata is missing ssm_layout") + layout = SSMShardLayout(**raw_layout) + self._validate_peer(meta, blocks, dst_block_ids) + peer_width = ( + layout.conv_dim_local if state_kind == "conv" else layout.nheads_local + ) + if ( + meta.get("heads_per_partition") != peer_width + or int(meta["num_outer"]) != layout.num_layers + or int(meta["blocks_axis"]) != 1 + ): + raise ValueError( + f"peer {state_kind} geometry does not match its SSM layout" + ) + if layout.global_rank in peers_by_rank: + raise ValueError( + f"duplicate source global_rank={layout.global_rank} " "in SSM metadata" + ) + sources.append(layout) + peers_by_rank[layout.global_rank] = (meta, blocks) + if not sources: + raise ValueError("SSM handoff contains no source peer metadata") + + transfers = [ + transfer + for transfer in plan_ssm_reshard(sources, [self._ssm_layout]) + if transfer.is_conv == (state_kind == "conv") + ] + for layer in range(self._ssm_layout.num_layers): + intervals = sorted( + (transfer.dst_lo, transfer.dst_hi) + for transfer in transfers + if transfer.dst_layer == layer + ) + if not intervals or intervals[0][0] != 0 or intervals[-1][1] != width: + raise ValueError(f"incomplete SSM {state_kind} coverage for layer {layer}") + if any(a[1] != b[0] for a, b in zip(intervals, intervals[1:])): + raise ValueError( + f"non-contiguous SSM {state_kind} coverage for layer {layer}" + ) + + for transfer in transfers: + meta, blocks = peers_by_rank[transfer.src_rank] + xfer, ctx = self._begin_transfer( + meta, + blocks, + dst_block_ids, + transfer.src_layer, + transfer.dst_layer, + 1, + transfer.src_lo, + transfer.dst_lo, + transfer.src_hi - transfer.src_lo, + ) + xfers.append(xfer) + contexts.append(ctx) + elif self._layout is not None: + sources = [] + peers_by_rank = {} + for meta, blocks in transfer_peer_records(peer_meta, src_block_ids): + layout = self._kv_layout_from_meta(meta) + self._validate_peer(meta, blocks, dst_block_ids) + if meta.get("heads_per_partition") != layout.local_num_heads(): + raise ValueError("peer heads_per_partition does not match its KV layout") + if int(meta["num_outer"]) % layout.local_num_layers(): + raise ValueError("peer num_outer is not divisible by its local layer count") + if layout.global_rank in peers_by_rank: + raise ValueError( + f"duplicate source global_rank={layout.global_rank} in KV metadata" + ) + sources.append(layout) + peers_by_rank[layout.global_rank] = (meta, blocks, layout) + if not sources: + raise ValueError("KV handoff contains no source peer metadata") + + local_planes = self._num_outer // self._layout.local_num_layers() + for transfer in plan_kv_reshard(sources, [self._layout]): + meta, blocks, source_layout = peers_by_rank[transfer.src_rank] + source_planes = int(meta["num_outer"]) // source_layout.local_num_layers() + if source_planes != local_planes: + raise ValueError( + f"outer-plane mismatch peer={source_planes} local={local_planes}" + ) + + src_layers = transfer.src_layer_slice(source_layout) + dst_layers = transfer.dst_layer_slice(self._layout) + src_heads = transfer.src_head_slice(source_layout) + dst_heads = transfer.dst_head_slice(self._layout) + layer_count = src_layers.stop - src_layers.start + head_count = src_heads.stop - src_heads.start + full_heads = ( + src_heads.start == 0 + and src_heads.stop == source_layout.local_num_heads() + and dst_heads.start == 0 + and dst_heads.stop == self._layout.local_num_heads() + and int(meta["bytes_per_slice"]) == self._bytes_per_slice + ) + full_layers = ( + src_layers.start == 0 + and src_layers.stop == source_layout.local_num_layers() + and dst_layers.start == 0 + and dst_layers.stop == self._layout.local_num_layers() + ) + if full_heads and full_layers: + xfer, ctx = self._begin_transfer( + meta, blocks, dst_block_ids, 0, 0, self._num_outer + ) + xfers.append(xfer) + contexts.append(ctx) + continue + if not full_heads and (int(meta["blocks_axis"]) != 2 or self._blocks_axis != 2): + raise NotImplementedError( + "KV head resharding requires the [2, L, B, T, H, d] layout" + ) + + for plane in range(local_planes): + xfer, ctx = self._begin_transfer( + meta, + blocks, + dst_block_ids, + plane * source_layout.local_num_layers() + src_layers.start, + plane * self._layout.local_num_layers() + dst_layers.start, + layer_count, + 0 if full_heads else src_heads.start, + 0 if full_heads else dst_heads.start, + 0 if full_heads else head_count, + ) + xfers.append(xfer) + contexts.append(ctx) + else: + records = transfer_peer_records(peer_meta, src_block_ids) + if len(records) != 1: + raise ValueError("matched-layout transfer requires exactly one source peer") + meta, blocks = records[0] + self._validate_peer(meta, blocks, dst_block_ids, matched_layout=True) + xfer, ctx = self._begin_transfer(meta, blocks, dst_block_ids, 0, 0, self._num_outer) + xfers.append(xfer) + contexts.append(ctx) + except Exception as exc: + if xfers: + cleanup = NixlPullHandle( + agent=self._agent, xfers=xfers, contexts=contexts, submitted_at=submitted_at + ) + try: + cleanup.wait() + except TimeoutError: + # Tell the owner not to recycle the destination storage while + # an already-submitted transfer may still write to it. + setattr(exc, "transfer_destinations_safe", False) + except Exception: + # Transfer errors are reported only after every submitted + # transfer has reached a terminal state. + pass + raise + return NixlPullHandle( + agent=self._agent, xfers=xfers, contexts=contexts, submitted_at=submitted_at + ) + + def _begin_transfer( + self, + peer_meta: Dict[str, Any], + src_block_ids: List[int], + dst_block_ids: List[int], + src_o_start: int, + dst_o_start: int, + n_outer: int, + src_h0: int = 0, + dst_h0: int = 0, + n_heads: int = 0, + ) -> tuple[Any, str]: + """Submit one full-slice or head-fragment NIXL transfer.""" + pm = peer_meta + peer_base = pm["base_addr"] + peer_device_id = pm.get("device_id", 0) + peer_bps = pm["bytes_per_slice"] + peer_os = pm["outer_stride_bytes"] + peer_id = self._ensure_peer_registered(pm) + + bps = self._bytes_per_slice + local_os = self._outer_stride_bytes + + src_tuples: List[Any] = [] + dst_tuples: List[Any] = [] + + if n_heads == 0: + # One descriptor per block and outer slice. + for src_b, dst_b in zip(src_block_ids, dst_block_ids): + for i in range(n_outer): + src_o = src_o_start + i + dst_o = dst_o_start + i + src_tuples.append( + (peer_base + src_o * peer_os + src_b * peer_bps, peer_bps, peer_device_id) + ) + dst_tuples.append( + (self._buf_ptr + dst_o * local_os + dst_b * bps, bps, self._device_id) + ) + ctx = ( + f"matched peer={peer_id} outer[{src_o_start}:+{n_outer}] " + f"blocks={len(src_block_ids)}" + ) + else: + # Head sub-range copy: one descriptor per token. + assert self._head_dim is not None + assert self._heads_per_partition is not None + assert self._tokens_per_block is not None + d_bytes = self._head_dim * self._element_size + local_token_stride = self._heads_per_partition * d_bytes + peer_token_stride = pm["heads_per_partition"] * d_bytes + T = self._tokens_per_block + frag_bytes = n_heads * d_bytes + src_h_off = src_h0 * d_bytes + dst_h_off = dst_h0 * d_bytes + + for src_b, dst_b in zip(src_block_ids, dst_block_ids): + for i in range(n_outer): + src_o = src_o_start + i + dst_o = dst_o_start + i + src_slice = peer_base + src_o * peer_os + src_b * peer_bps + src_h_off + dst_slice = self._buf_ptr + dst_o * local_os + dst_b * bps + dst_h_off + for t in range(T): + src_tuples.append( + (src_slice + t * peer_token_stride, frag_bytes, peer_device_id) + ) + dst_tuples.append( + (dst_slice + t * local_token_stride, frag_bytes, self._device_id) + ) + ctx = ( + f"reshard peer={peer_id} outer[{src_o_start}:+{n_outer}] " + f"heads[{src_h0}:+{n_heads}] blocks={len(src_block_ids)}" + ) + + src_descs = self._agent.get_xfer_descs(src_tuples, mem_type="VRAM") + dst_descs = self._agent.get_xfer_descs(dst_tuples, mem_type="VRAM") + # READ pulls remote -> local. Signature is (op, local, remote, peer). + xfer = self._agent.initialize_xfer("READ", dst_descs, src_descs, peer_id) + try: + self._agent.transfer(xfer) + except Exception as exc: + # The transport may have accepted the operation before surfacing an + # error, so its destination cannot be proven safe for immediate reuse. + setattr(exc, "transfer_destinations_safe", False) + raise + return xfer, ctx + + def close(self) -> None: + """Release the registration and agent.""" + if self._agent is None: + return + try: + self._agent.deregister_memory(self._reg_handle) + except Exception: # noqa: BLE001 - shutdown path + logger.exception("NixlTransferBackend: deregister_memory failed") + self._agent = None diff --git a/megatron/core/inference/disaggregation/utils.py b/megatron/core/inference/disaggregation/utils.py index 9b5e153b443..779207e40f6 100644 --- a/megatron/core/inference/disaggregation/utils.py +++ b/megatron/core/inference/disaggregation/utils.py @@ -4,7 +4,7 @@ from __future__ import annotations -from typing import Optional, Tuple +from typing import Any, List, Optional, Tuple def intersect(a: Tuple[int, int], b: Tuple[int, int]) -> Optional[Tuple[int, int]]: @@ -14,7 +14,7 @@ def intersect(a: Tuple[int, int], b: Tuple[int, int]) -> Optional[Tuple[int, int def transfers_for_src(plan, src_rank): - """Transfers in ``plan`` originating from ``src_rank`` (any KV/Mamba + """Transfers in ``plan`` originating from ``src_rank`` (any KV/SSM reshard transfer -- both expose a ``src_rank`` field).""" return [t for t in plan if t.src_rank == src_rank] @@ -22,3 +22,36 @@ def transfers_for_src(plan, src_rank): def transfers_for_dst(plan, dst_rank): """Transfers in ``plan`` destined for ``dst_rank``.""" return [t for t in plan if t.dst_rank == dst_rank] + + +def transfer_peer_records(peer_meta: Any, src_block_ids: List[int]) -> List[Tuple[dict, List[int]]]: + """Normalize flat/TP/PP transfer metadata into peer/block records.""" + + def append_metas(raw_metas: Any, default_blocks: List[int]) -> None: + metas = raw_metas if isinstance(raw_metas, list) else [raw_metas] + for meta in metas: + if not isinstance(meta, dict): + raise ValueError("transfer peer metadata entries must be dictionaries") + blocks = meta.get("block_ids", default_blocks) + records.append((meta, [int(block) for block in blocks])) + + records: List[Tuple[dict, List[int]]] = [] + if isinstance(peer_meta, dict) and "pp_metas" in peer_meta: + for entry in peer_meta["pp_metas"]: + raw_metas = entry.get("tp_metas", entry) + blocks = [int(block) for block in entry.get("block_ids", [])] + append_metas(raw_metas, blocks) + return records + + if isinstance(peer_meta, dict) and "tp_metas" in peer_meta: + peer_meta = peer_meta["tp_metas"] + blocks = [int(block) for block in src_block_ids] + append_metas(peer_meta, blocks) + return records + + +def transfer_block_count(peer_meta: Any, src_block_ids: List[int]) -> int: + """Return the sequence-block count represented by transfer metadata.""" + + records = transfer_peer_records(peer_meta, src_block_ids) + return len(records[0][1]) if records else 0 diff --git a/megatron/core/inference/engines/async_zmq_communicator.py b/megatron/core/inference/engines/async_zmq_communicator.py index aa13f659d40..bba7508e08b 100644 --- a/megatron/core/inference/engines/async_zmq_communicator.py +++ b/megatron/core/inference/engines/async_zmq_communicator.py @@ -41,6 +41,12 @@ def __init__( hostname (str | None): Hostname or IP address to use for ZMQ socket binding. If None, defaults to socket.gethostname(). """ + # Normalize None to the default (world) group. get_rank/get_world_size + # already treat None this way, but get_process_group_ranks below does + # not accept None, so resolve it once here for all three calls. + if process_group is None: + process_group = dist.group.WORLD + self.rank = dist.get_rank(process_group) self.world_size = dist.get_world_size(process_group) self.is_leader = self.rank == 0 diff --git a/megatron/core/inference/engines/dynamic_engine.py b/megatron/core/inference/engines/dynamic_engine.py index 944e8f28c46..833b65dd15f 100644 --- a/megatron/core/inference/engines/dynamic_engine.py +++ b/megatron/core/inference/engines/dynamic_engine.py @@ -19,6 +19,10 @@ import torch from torch import Tensor +from megatron.core.inference.batch_dimensions_utils import ( + CUDAGraphBatchDimensionBuilder, + InferenceBatchDimensions, +) from megatron.core.inference.config import AsyncScheduleMode, KVCacheManagementMode from megatron.core.inference.contexts.dynamic_context import ( BlockOverflowError, @@ -40,6 +44,8 @@ ) from megatron.core.inference.sampling_params import SamplingParams from megatron.core.inference.text_generation_controllers.text_generation_controller import ( + DecodeOnly, + DynamicBatchControllerStepResult, TextGenerationController, ) from megatron.core.inference.utils import Counter, InferenceMode, await_process_call @@ -142,6 +148,41 @@ def format_mem_bytes(mem_bytes): return "%d bytes" % mem_bytes +def _get_decode_only_log_state( + mode: AsyncScheduleMode, decode_only: DecodeOnly +) -> Tuple[str, Optional[bool]]: + """Build the console transition label and color state for one inference step. + + Args: + mode (AsyncScheduleMode): Active scheduling mode. + decode_only (DecodeOnly): Decode-only state for the consumed and launched forwards. + + Returns: + Tuple[str, Optional[bool]]: Current step label, including the previous + step when it differs, and whether to use decode coloring. + """ + if mode == AsyncScheduleMode.LEGACY: + is_decode_only = bool(decode_only) + return ("decode" if is_decode_only else "non-decode"), is_decode_only + + current_decode_only = ( + decode_only.launched if decode_only.launched is not None else decode_only.consumed + ) + if current_decode_only is None: + return "idle", None + + step_type = "decode" if current_decode_only else "non-decode" + if ( + decode_only.consumed is not None + and decode_only.launched is not None + and decode_only.consumed != decode_only.launched + ): + previous_step_type = "decode" if decode_only.consumed else "non-decode" + step_type = f"{step_type} (prev: {previous_step_type})" + + return step_type, current_decode_only + + def _cuda_graph_mempool_bytes() -> Tuple[int, int]: """Return (reserved, allocated) bytes belonging to the global CUDA graph mempool. @@ -246,6 +287,7 @@ def __init__(self, controller: TextGenerationController, context: DynamicInferen self.track_paused_request_events = inference_config.track_paused_request_events self.track_generated_token_events = inference_config.track_generated_token_events self.enable_chunked_prefill = inference_config.enable_chunked_prefill + self.cuda_graph_all_prefills = inference_config.cuda_graph_all_prefills self.metrics_writer = inference_config.metrics_writer self.logging_step_interval = inference_config.logging_step_interval self.unified_memory_level = inference_config.unified_memory_level @@ -324,6 +366,7 @@ def reset(self) -> None: self.capture_stats = None # Runtime state. + self.decode_only = DecodeOnly(consumed=None, launched=None) self._loop = get_asyncio_loop(getattr(self, "_loop", None)) self._cond = asyncio.Condition() self._state_events = {k: asyncio.Event() for k in self._STATE_EVENTS} @@ -347,6 +390,8 @@ def reset(self) -> None: # Prefix caching tracking. self._prefix_cache_hits = 0 self._prefix_cache_blocks_matched = 0 + self._prefill_tokens_computed = 0 + self._prefill_tokens_skipped = 0 self._prefix_coordination_waits = 0 # Coordinator state. @@ -434,7 +479,7 @@ def create_cuda_graphs(self, reset_context: bool = True): if HAVE_TQDM: tbar = tqdm(tbar, total=len(context.cuda_graph_batch_dimensions_list)) for tbar_idx, cuda_graph_batch_dimension in tbar: - input_ids, position_ids = self.controller._dynamic_step_context_init( + input_ids, position_ids, _ = self.controller._dynamic_step_context_init( construct_graph_dimensions=cuda_graph_batch_dimension ) # Progress. @@ -986,55 +1031,35 @@ def get_request(self, request_id: int) -> DynamicInferenceRequest: return self.requests[request_id].record[-1] def _validate_async_sched_support_for_config(self) -> None: - """Validate config-level restrictions for serial async scheduling. + """Validate config-level restrictions for async scheduling. - Raises if the config does not support serial async scheduling. + Raises if the config does not support async scheduling. """ - if self.context.config.async_sched_mode != AsyncScheduleMode.SERIAL: + mode = self.context.config.async_sched_mode + if mode == AsyncScheduleMode.LEGACY: return + if mode != AsyncScheduleMode.ASYNC: + raise AssertionError(f"Unexpected async scheduling mode: {mode}") model_config = self.controller.inference_wrapped_model.model.config - if self.num_speculative_tokens > 0: - raise ValueError("Async scheduling does not support speculative tokens.") - if self.context.is_hybrid_model: - raise ValueError("Async scheduling does not support hybrid/Mamba models.") - if self.context.enable_prefix_caching: - raise ValueError("Async scheduling does not support prefix caching.") - if not self.materialize_only_last_token_logits: - raise ValueError("Async scheduling requires materialize_only_last_token_logits=True.") - if model_config.expert_model_parallel_size > 1: - raise ValueError("Async scheduling does not support expert parallelism.") - if model_config.num_moe_experts is not None: - raise ValueError("Async scheduling does not support MoE models.") + if self.num_speculative_tokens > self.controller.num_mtp_depths: + raise ValueError("Async scheduling requires one MTP depth per speculative token.") if model_config.moe_enable_routing_replay: raise ValueError("Async scheduling does not support routing replay.") - def _validate_async_sched_support_for_request(self, request: DynamicInferenceRequest) -> None: - """Validate request-level restrictions for serial async scheduling. - - Args: - request (DynamicInferenceRequest): Request being added to the engine. - """ - if self.context.config.async_sched_mode != AsyncScheduleMode.SERIAL: - return - - sampling_params = request.sampling_params - if sampling_params.top_k != 1 or sampling_params.top_p != 0.0: - raise ValueError( - "Async scheduling only supports greedy sampling " - "(SamplingParams.top_k == 1 and top_p == 0.0)." - ) - if sampling_params.return_log_probs or sampling_params.top_n_logprobs > 0: - raise ValueError("Async scheduling does not support log probabilities.") - if sampling_params.stop_words: - raise ValueError("Async scheduling does not support stop words.") - def _add_request( self, request: DynamicInferenceRequest ) -> asyncio.Future[DynamicInferenceRequest]: + """Add a request to the engine. + + Args: + request (DynamicInferenceRequest): Request to add. + + Returns: + asyncio.Future[DynamicInferenceRequest]: Future completed when the request finishes. + """ request_id = request.request_id - self._validate_async_sched_support_for_request(request) # Add request to self.requests. If the engine has previously been # suspended, then the request may already exist. @@ -1206,6 +1231,7 @@ def post_process_requests( sample: torch.Tensor, accepted_tokens: torch.Tensor, log_probs: torch.Tensor, + consumed_chunked_prefill_request_id: int, top_n_logprobs: Optional[Dict[int, List[Tuple[torch.Tensor, torch.Tensor]]]] = None, pre_fwd_active_token_count: Optional[int] = None, pre_fwd_step_count: Optional[int] = None, @@ -1222,8 +1248,13 @@ def post_process_requests( sample: Tensor: The newly generated token for each request accepted_tokens: Tensor: The additional accepted tokens for each request log_probs: (List): Log probs for each request + consumed_chunked_prefill_request_id (int): Chunked-prefill request ID + associated with the consumed forward, or -1 if it had no partial chunk. top_n_logprobs: (Dict): Top-n log probs for each request. Maps request_idx to list of (top_n_logprobs, top_n_indices) tuples. + pre_fwd_active_token_count (Optional[int]): Active token count for the + consumed forward. + pre_fwd_step_count (Optional[int]): Step count for the consumed forward. finished_routing_block_ids: (Dict[int, List[int]]): Block IDs for finished requests, saved before update_requests released them. Used for per-block routing reconstruction. @@ -1280,7 +1311,7 @@ def post_process_requests( num_stop_word_trim = 0 is_prefill = len(request.generated_tokens) == 0 - if request_id != self.context.chunked_prefill_request_id: + if request_id != consumed_chunked_prefill_request_id: # Skip appending token for requests being finished due to stop words # (they already have their final token from the previous step) # If the request already has more tokens, then we only append as much as is necessary @@ -1434,7 +1465,7 @@ def post_process_requests( if not request.generated_log_probs: request.generated_log_probs = [] - is_chunked_prefill = request_id == self.context.chunked_prefill_request_id + is_chunked_prefill = request_id == consumed_chunked_prefill_request_id is_prefill = len(request.generated_log_probs) == 0 if request.sampling_params.skip_prompt_log_probs: @@ -1600,24 +1631,8 @@ def get_prefix_coordination_metrics(self) -> dict: """ return {"waits": self._prefix_coordination_waits} - def _find_mamba_match_count(self, req: DynamicInferenceRequest) -> int: - """Find farthest block with cached Mamba state by iterating from the end. - - Not all blocks have Mamba state cached in mamba_hash_to_block_id, - only divergence and last-aligned blocks do. Iterating from the end - finds the farthest block with cached state, which is the only one - needed for restore since Mamba state is cumulative. - """ - if not req.precomputed_block_hashes: - return 0 - mamba_map = self.context.mamba_slot_allocator.hash_to_block_id - for i in range(len(req.precomputed_block_hashes) - 1, -1, -1): - if req.precomputed_block_hashes[i] in mamba_map: - return i + 1 - return 0 - - def schedule_waiting_requests(self): - """Tries to schedule any requests in the waiting pool.""" + def schedule_waiting_requests(self) -> None: + """Try to schedule requests from the waiting pool.""" # Keep track of which requests get scheduled. waiting_before = set(self.waiting_request_ids) if self.enable_chunked_prefill: @@ -1633,16 +1648,73 @@ def schedule_waiting_requests(self): if req.kv_cache_epoch is None: req.kv_cache_epoch = [(0, self._generation_epoch)] - def schedule_non_chunked_prefill(self): + def _can_schedule_non_chunked_prefill(self, req, *, record_cg_wait: bool) -> bool: + """Return whether the queue-head request can be admitted now. + + Args: + req: Queue-head inference request. + record_cg_wait (bool): Whether a CUDA-graph miss should update the + request's wait counter. + + Returns: + bool: Whether all request, token, KV-cache, and CUDA-graph checks pass. """ - Perform the same original scheduling logic for non-chunked runs + if not all(self.context.check_availability(req)): + return False + + if not self._cg_admission_gating_active(): + return True + + candidate = InferenceBatchDimensions( + token_count=self.context.active_token_count + len(req.remaining_prompt_tokens), + prefill_req_count=self.context.num_prefill_requests + 1, + decode_req_count=self.context.num_decode_requests, + ) + if record_cg_wait: + return self._cg_admission_check(req, candidate) + return self._matches_cg_admission(candidate) + + def _can_schedule_chunked_prefill(self, req) -> bool: + """Return whether the queue-head request can admit at least one prompt token. + + Args: + req: Queue-head inference request. + + Returns: + bool: Whether request, token, and KV-cache capacity permit a chunk. """ - prefix_caching_enabled = self.context.enable_prefix_caching - mamba_caching_enabled = ( - prefix_caching_enabled - and self.context.is_hybrid_model - and self.context.mamba_slot_allocator is not None + request_can_be_added, _, kv_cache_available = self.context.check_availability(req) + is_continuing_chunk = self.context.chunked_prefill_request_id == req.request_id + token_capacity_available = self.context.active_token_count < self.context.max_tokens + return ( + (is_continuing_chunk or request_can_be_added) + and kv_cache_available + and token_capacity_available ) + + def _should_run_async_sched_overlap(self) -> bool: + """Return whether this step should use overlap ordering. + + Returns: + bool: Whether the next step can use overlap ordering. + """ + # No-overlap also handles the first decode-only forward after prefill: + # pending prefill output must be resolved before preparing its decode rows. + # Paused requests and insufficient KV capacity likewise require complete + # lifecycle bookkeeping before preparing the next batch. + if not self.context.can_prepare_requests(): + return False + if not self.waiting_request_ids: + return True + + req = self.get_request(self.waiting_request_ids[0]) + if self.enable_chunked_prefill: + return not self._can_schedule_chunked_prefill(req) + return not self._can_schedule_non_chunked_prefill(req, record_cg_wait=False) + + def schedule_non_chunked_prefill(self) -> None: + """Schedule non-chunked prefill requests.""" + prefix_caching_enabled = self.context.enable_prefix_caching if prefix_caching_enabled: pending_block_hashes = set() pending_request_ids = [] @@ -1661,14 +1733,7 @@ def schedule_non_chunked_prefill(self): pending_request_ids.append(self.waiting_request_ids.popleft()) continue - # Find Mamba prefix match before check_availability (sets skip count) - if mamba_caching_enabled: - req._mamba_num_matched_blocks = self._find_mamba_match_count(req) - - request_can_be_added, request_tokens_can_be_added, kv_cache_available = ( - self.context.check_availability(req) - ) - if request_can_be_added and request_tokens_can_be_added and kv_cache_available: + if self._can_schedule_non_chunked_prefill(req, record_cg_wait=True): # Add these hashes to pending. if prefix_caching_enabled: for block_hash in req.precomputed_block_hashes: @@ -1688,6 +1753,108 @@ def schedule_non_chunked_prefill(self): if prefix_caching_enabled and pending_request_ids: self.waiting_request_ids.extendleft(reversed(pending_request_ids)) + def _cg_admission_gating_active(self) -> bool: + """Cudagraph-aware admission gating is active when --inference-cuda-graph-all-prefills + is set, the engine has prefill/mixed CGs, and the batch-dim list is populated. + + All are required so legacy tests that exercise the scheduler without intending to run on + captured graphs are unaffected. Gating is opt-in via `cuda_graph_all_prefills`. + """ + return ( + self.cuda_graph_all_prefills + and self.context.use_cuda_graphs_for_non_decode_steps + and bool(self.context.cuda_graph_batch_dimensions_list) + ) + + def _find_cg_chunk_size(self, max_chunk_tokens: int) -> Optional[int]: + """Return the largest chunk size <= max_chunk_tokens where batch matches a captured graph, + or None if no graph covers any chunk in the budget. + + Walks the captured-CG list (sorted descending by token_count) and returns the first chunk + that falls within budget and produces an applicable batch_dim under the engine's matching + mode (strict for hybrid models). Callers must explicitly handle the None case by deferring + the admission rather than scheduling eagerly. + """ + active_tok = self.context.active_token_count + active_p = self.context.num_prefill_requests + active_d = self.context.num_decode_requests + strict = self.context.is_hybrid_model + + for cg in self.context.cuda_graph_batch_dimensions_list: + chunk = cg.token_count - active_tok + if chunk < 1: + continue + if chunk > max_chunk_tokens: + continue + candidate = InferenceBatchDimensions( + token_count=cg.token_count, + prefill_req_count=active_p + 1, + decode_req_count=active_d, + ) + # candidate.token_count == cg.token_count, so the token-dimension check inside + # is_applicable_for_batch_dim is always True here; this call filters on P/D compatibility only. + if cg.is_applicable_for_batch_dim(candidate, strict=strict): + return chunk + + return None + + def _register_cg_wait(self, req) -> None: + """Track a deferred admission attempt and throw a starvation warning at the threshold. + + Decode is bounded by the number of decode steps. + Persistent waits past `_cg_admission_warn_after` consecutive steps signal a problem. + """ + req.cg_wait_iters += 1 + if req.cg_wait_iters % self._cg_admission_warn_after == 0: + logging.warning( + "request %d has been deferred by CG-aware admission for %d steps — " + "possible starvation (strict=%s, active P=%d D=%d tok=%d)", + req.request_id, + req.cg_wait_iters, + self.context.is_hybrid_model, + self.context.num_prefill_requests, + self.context.num_decode_requests, + self.context.active_token_count, + ) + + def _cg_admission_check(self, req, candidate: InferenceBatchDimensions) -> bool: + """Return True if the candidate batch shape matches a captured cudagraph. + + On miss, registers a wait + warning via `_register_cg_wait`. On hit, resets the counter. + Caller is responsible for breaking the scheduler loop on False. + Passes match_ep_token_counts=False so this local admission probe doesn't force a per-attempt + NCCL all-reduce — the step-time matcher does its own EP sync. + + Args: + req: Request whose CUDA-graph wait state should be updated. + candidate (InferenceBatchDimensions): Candidate batch after admission. + + Returns: + bool: Whether a compatible captured graph exists. + """ + if self._matches_cg_admission(candidate): + req.cg_wait_iters = 0 + return True + self._register_cg_wait(req) + return False + + def _matches_cg_admission(self, candidate: InferenceBatchDimensions) -> bool: + """Return whether a candidate batch matches a captured CUDA graph. + + Args: + candidate (InferenceBatchDimensions): Candidate batch after admission. + + Returns: + bool: Whether a compatible captured graph exists. + """ + matched = CUDAGraphBatchDimensionBuilder.match_graph_config( + real_batch_dim=candidate, + cuda_graph_batch_dimensions_list=self.context.cuda_graph_batch_dimensions_list, + strict=self.context.is_hybrid_model, + match_ep_token_counts=False, + ) + return matched is not None + def schedule_chunked_prefill(self): """ This function schedules chunked prefill requests. @@ -1704,11 +1871,6 @@ def schedule_chunked_prefill(self): - For each request, remaining_prompt_tokens holds the **unprefilled** prompt tokens """ prefix_caching_enabled = self.context.enable_prefix_caching - mamba_caching_enabled = ( - prefix_caching_enabled - and self.context.is_hybrid_model - and self.context.mamba_slot_allocator is not None - ) if prefix_caching_enabled: pending_block_hashes = set() pending_request_ids = [] @@ -1736,29 +1898,109 @@ def schedule_chunked_prefill(self): ) continue - # Find Mamba prefix match for non-continuing requests - if mamba_caching_enabled and not is_continuing_chunked_prefill: - req._mamba_num_matched_blocks = self._find_mamba_match_count(req) - # Use remaining prompt tokens for scheduling decisions remaining_len = len(req.remaining_prompt_tokens) - token_fully_can_be_added = ( - self.context.active_token_count + remaining_len <= self.context.max_tokens - ) - token_partially_can_be_added = self.context.active_token_count < self.context.max_tokens - request_can_be_added, _, kv_cache_available = self.context.check_availability(req) - request_can_be_added = is_continuing_chunked_prefill or request_can_be_added - - if request_can_be_added and kv_cache_available: - if token_fully_can_be_added: - # Add these hashes to pending. - if prefix_caching_enabled: - for block_hash in req.precomputed_block_hashes: - if ( - block_hash - not in self.context.kv_block_allocator.kv_hash_to_block_id - ): - pending_block_hashes.add(block_hash) + + if self._can_schedule_chunked_prefill(req): + # How many tokens we can admit this step. + token_budget = self.context.max_tokens - self.context.active_token_count + + # Prefix-cache skip: on a request's first chunk, the tokens covered + # by a cached prefix are reused rather than recomputed, so they do + # NOT consume the compute budget. Extend this chunk's SPAN to cover + # the entire skippable prefix plus up to `token_budget` newly computed + # tokens. Without this the span is capped at the budget, forcing the + # rest of a long cached prefix to be re-prefilled over many chunks + # (latency then scales with prompt length instead of the delta). + # add_request() only computes `effective = span - skip` tokens. + prefix_skip = 0 + if prefix_caching_enabled and not is_continuing_chunked_prefill: + (_, _, _, _, prefix_skip, _) = self.context._compute_prefix_match( + req, remaining_len + ) + prefix_skip = min(prefix_skip, remaining_len - 1) # keep >=1 token to run + + computed_budget = min(remaining_len - prefix_skip, token_budget) + + # Skip CG gating for the continuation of an in-flight chunked prefill: + # the request is already mid-flight, deferring it would deadlock progress. + if self._cg_admission_gating_active() and not is_continuing_chunked_prefill: + # Snap the COMPUTED chunk size to the largest captured-CG boundary + # within budget (skipped tokens don't affect the CG batch shape). + # Fall back to eager (computed_budget) if no CG shape covers it. + snapped_chunk = self._find_cg_chunk_size(computed_budget) + computed_chunk = snapped_chunk if snapped_chunk is not None else computed_budget + req.cg_wait_iters = 0 + else: + computed_chunk = computed_budget + + prefill_chunk_length = prefix_skip + computed_chunk + + # Mamba prefix caching: keep chunk boundaries block-aligned. + # compute_and_store_offsets() records a recurrent-state snapshot at a + # KV-block boundary only when that boundary lands on a multiple of the + # SSM chunk size measured FROM the start of the current prefill chunk + # (it filters on `offset % mamba_chunk_size == 0`, where the chunk start + # equals `finished_chunk_token_count` on continuation chunks). Block + # boundaries are multiples of `block_size_tokens` (itself a multiple of + # the SSM chunk size), so the filter only passes when + # `finished_chunk_token_count` is block-aligned. If a chunk ends at an + # arbitrary token offset, every candidate boundary in the following + # chunks becomes unrecordable and the last-block snapshot that lets a + # future request skip prefill is silently dropped. Stop a partial + # (non-final) chunk short at the nearest lower block boundary so the + # running `finished_chunk_token_count` stays block-aligned. + if ( + self.context.is_hybrid_model + and self.context.mamba_slot_allocator is not None + and prefill_chunk_length < remaining_len + ): + block_size = self.context.block_size_tokens + chunk_end = req.finished_chunk_token_count + prefill_chunk_length + aligned_end = (chunk_end // block_size) * block_size + aligned_chunk_length = aligned_end - req.finished_chunk_token_count + # Only snap down when the aligned chunk still computes at least one + # token beyond the skipped prefix (a chunk whose budget is smaller + # than a block cannot be block-aligned; leave it unchanged). + if aligned_chunk_length > prefix_skip: + prefill_chunk_length = aligned_chunk_length + + # Flash-attn guard: if this chunk would leave exactly 1 token for the + # final chunk, reduce by 1 (or defer if we only have 1 computed token). + # See https://github.com/Dao-AILab/flash-attention/issues/1537 + # The -1 is safe after CG snapping: is_applicable_for_batch_dim matches on + # cg.token_count >= real.token_count, so the snapped CG still covers token_count-1. + if remaining_len - prefill_chunk_length == 1: + if computed_chunk > 1: + prefill_chunk_length -= 1 + else: + can_schedule = False + break + + # add_request recomputes the skip for this exact chunk and applies a + # ">= 2 computed tokens" clamp. When the chunk would compute fewer than + # 2 tokens (tight budget late in a batched step, or a prompt that is + # all-but-one cached) that clamp shrinks the skip and grows the computed + # count by up to one block, which can exceed the token budget + # (TokenOverflowError). Only then re-derive the exact effective length + # add_request will use and defer on overflow (a later full-budget step + # admits the request). For >= 2 computed tokens add_request computes + # exactly this chunk, which already fits the budget. + if prefix_skip > 0 and (prefill_chunk_length - prefix_skip) < 2: + (_, _, _, _, _, actual_effective) = self.context._compute_prefix_match( + req, prefill_chunk_length + ) + if self.context.active_token_count + actual_effective > self.context.max_tokens: + can_schedule = False + break + + # Add hashes to pending set (prefix-caching bookkeeping). + if prefix_caching_enabled: + for block_hash in req.precomputed_block_hashes: + if block_hash not in self.context.kv_block_allocator.kv_hash_to_block_id: + pending_block_hashes.add(block_hash) + + if prefill_chunk_length >= remaining_len: self.context.chunked_prefill_request_id = -1 self.context.add_request(req) self._loop.call_soon_threadsafe( @@ -1766,35 +2008,10 @@ def schedule_chunked_prefill(self): ) req.remaining_prompt_tokens = req.remaining_prompt_tokens.new_empty(0) req.add_event_add_context() - # Fully scheduled, so we remove from waiting pool self.waiting_request_ids.popleft() - # Only this case we keep checking the rest of the waiting queue can_schedule = True - elif token_partially_can_be_added: - # Add these hashes to pending. - if prefix_caching_enabled: - for block_hash in req.precomputed_block_hashes: - if ( - block_hash - not in self.context.kv_block_allocator.kv_hash_to_block_id - ): - pending_block_hashes.add(block_hash) - prefill_chunk_length = self.context.max_tokens - self.context.active_token_count - - # If this chunk would leave exactly 1 token for the final chunk, reduce - # this chunk by 1 or skip scheduling so the final chunk has 2 tokens. - # This avoids the edge case where max_seqlen_q=1 which results in a bug - # with the Flash Attention kernel. - # See https://github.com/Dao-AILab/flash-attention/issues/1537 - if remaining_len - prefill_chunk_length == 1: - if prefill_chunk_length > 1: - prefill_chunk_length -= 1 - else: - # We only have space for 1 token, but remaining is 2. - # Delay scheduling to avoid leaving exactly 1 token for the final chunk. - can_schedule = False - break - + else: + # Partial admit: schedule this chunk and keep the request at the queue head. self.context.add_request(req, prefill_chunk_length=prefill_chunk_length) self._loop.call_soon_threadsafe( self._loop.create_task, self._notify_cond_for_new_request() @@ -1802,9 +2019,6 @@ def schedule_chunked_prefill(self): self.context.chunked_prefill_request_id = req.request_id req.remaining_prompt_tokens = req.remaining_prompt_tokens[prefill_chunk_length:] req.finished_chunk_token_count += prefill_chunk_length - # Still have tokens to prefill, so we break and keep the - # chunked prefill request at the head of the waiting queue - # Note that we do not need to continue check the queue, as the tokens are full # Prepend pending request ids to waiting queue. if prefix_caching_enabled and pending_request_ids: @@ -1816,15 +2030,15 @@ def schedule_chunked_prefill(self): else: self.waiting_request_ids.extendleft(reversed(pending_request_ids)) - async def async_forward(self) -> Tuple[Dict, Dict, float]: + async def async_forward(self) -> Tuple[Optional[Dict], Dict, float]: """Uses `asyncio` for continuous generation. Sleeps when no requests are available, until new requests have been added. Returns: A tuple comprised of: step_result (Optional[Dict]): The result of the step. - context_state (Dict): A tuple consisting of the state of the context. - is_decode_only, total/paused request count, active token count. + context_state (Dict): Decode-only state, total/paused request + count, and active token count. step_time (float): How long this step took. """ @@ -1832,8 +2046,22 @@ async def async_forward(self) -> Tuple[Dict, Dict, float]: if self.state in (EngineState.SUSPENDED, EngineState.SUSPENDING): raise EngineSuspendedError(self.context.step_count) - # schedule requests - self.schedule_waiting_requests() + mode = self.context.config.async_sched_mode + if mode == AsyncScheduleMode.LEGACY: + self.schedule_waiting_requests() + step_nvtx_range = "Decode" if self.context.num_prefill_requests == 0 else "Prefill" + controller_kwargs = {} + elif mode == AsyncScheduleMode.ASYNC: + run_async_overlap = self._should_run_async_sched_overlap() + step_nvtx_range = "AsyncOverlap" if run_async_overlap else "AsyncNoOverlap" + controller_kwargs = { + "run_async_overlap": run_async_overlap, + "schedule_waiting_requests": ( + None if run_async_overlap else self.schedule_waiting_requests + ), + } + else: + raise AssertionError(f"Unexpected async scheduling mode: {mode}") # The print block (async_bookkeep) and metrics block both fire on this # condition after step_count is incremented. Predict it up-front so we @@ -1844,10 +2072,8 @@ async def async_forward(self) -> Tuple[Dict, Dict, float]: and (self.context.step_count + 1) % self.logging_step_interval == 0 ) - is_decode_only = self.context.is_decode_only() if will_log_this_step: pre_step_context_state = { - "is_decode_only": is_decode_only, "max_requests": self.context.max_requests, "total_request_count": self.context.total_request_count, "paused_request_count": self.context.paused_request_count, @@ -1862,15 +2088,21 @@ async def async_forward(self) -> Tuple[Dict, Dict, float]: "active_token_count": self.context.active_token_count, "step_count": self.context.step_count, } + pre_step_context_state["chunked_prefill_request_id"] = ( + self.context.chunked_prefill_request_id + ) # Generate tokens. - nvtx_range_push("Prefill" if not is_decode_only else "Decode") - # TODO @TDE: Account for this line when overlapping forward and bookkeep. - self.is_decode_only = is_decode_only + nvtx_range_push(step_nvtx_range) if will_log_this_step: self.step_start_event.record() - result = await self.controller.async_generate_output_tokens_dynamic_batch() + controller_result: DynamicBatchControllerStepResult = ( + await self.controller.async_generate_output_tokens_dynamic_batch(**controller_kwargs) + ) + self.decode_only = controller_result.decode_only + pre_step_context_state["decode_only"] = self.decode_only + result = controller_result.output if will_log_this_step: self.step_end_event.record() self.step_end_event.synchronize() @@ -1880,7 +2112,7 @@ async def async_forward(self) -> Tuple[Dict, Dict, float]: self.context.step_count += 1 self.context.prefix_cache_lru_clock += 1 - nvtx_range_pop("Prefill" if not is_decode_only else "Decode") + nvtx_range_pop(step_nvtx_range) if will_log_this_step: kvcache_util_stats = ( @@ -1913,7 +2145,8 @@ async def async_bookkeep( Args: step_result (Optional[Dict]): The result of the step. - context_state (Dict): is_decode_only, total/paused request count, active token count. + context_state (Dict): Decode-only state, total/paused request count, + and active token count. step_time (float): How long this step took. Returns: @@ -1953,7 +2186,8 @@ async def async_bookkeep( sample, accepted_tokens, log_probs, - top_n_logprobs, + consumed_chunked_prefill_request_id=context_state["chunked_prefill_request_id"], + top_n_logprobs=top_n_logprobs, pre_fwd_active_token_count=context_state.get("active_token_count"), pre_fwd_step_count=context_state.get("step_count"), finished_routing_block_ids=finished_routing_block_ids, @@ -2015,8 +2249,12 @@ async def async_bookkeep( if self.context.enable_prefix_caching: self._prefix_cache_hits += self.context.prefix_cache_hits self._prefix_cache_blocks_matched += self.context.prefix_cache_blocks_matched + self._prefill_tokens_computed += self.context.prefix_cache_prefill_computed_tokens + self._prefill_tokens_skipped += self.context.prefix_cache_prefill_skipped_tokens self.context.prefix_cache_hits = 0 self.context.prefix_cache_blocks_matched = 0 + self.context.prefix_cache_prefill_computed_tokens = 0 + self.context.prefix_cache_prefill_skipped_tokens = 0 # Log KV cache utilization stats to W&B nvtx_range_push("wandb_logging") @@ -2081,7 +2319,10 @@ async def async_bookkeep( nvtx_range_push("cuda_memory_stats") mem = torch.cuda.memory_stats() nvtx_range_pop("cuda_memory_stats") - step_type = "decode" if context_state["is_decode_only"] else "non-decode" + decode_only = context_state["decode_only"] + step_type, color_decode_only = _get_decode_only_log_state( + self.context.config.async_sched_mode, decode_only + ) output_str = ( "* rank %d | step %d | %s ... time: %.3f ms%s ... " "reqs: a %d/%d, p %d, w %d, f %d, e %d ... " @@ -2144,7 +2385,37 @@ async def async_bookkeep( self._prefix_cache_hits, self._prefix_cache_blocks_matched, ) - if context_state["is_decode_only"]: + if self.context.enable_prefix_caching: + # Prefill compute actually saved by prefix caching (cumulative). + # computed = prompt tokens run through the model; skipped = prompt + # tokens whose prefill was reused from cache. If skipped% stays high + # while per-step latency grows, the growth is attention over the + # growing KV context, NOT re-prefilling skipped tokens. + _computed = self._prefill_tokens_computed + _skipped = self._prefill_tokens_skipped + _total = _computed + _skipped + output_str += " ... prefill (cumul): computed %d, skipped %d (%.1f%% skipped)" % ( + _computed, + _skipped, + (100.0 * _skipped / _total) if _total > 0 else 0.0, + ) + # Current cache occupancy (utilization). A Mamba durable-slot count + # near its max indicates the cache is saturating and will start + # LRU-evicting cached prefixes (hybrid models can only skip prefill + # where Mamba state is still cached). + kv_alloc = self.context.kv_block_allocator + output_str += " ... prefix cache util: KV %d/%d blocks cached (%d evictable)" % ( + len(kv_alloc.kv_hash_to_block_id), + kv_alloc.total_count, + int(kv_alloc.get_evictable_block_count()), + ) + msa = self.context.mamba_slot_allocator + if msa is not None: + output_str += ", mamba %d/%d durable slots" % ( + msa.max_slots - msa.free_count, + msa.max_slots, + ) + if color_decode_only: output_str = f"\033[94m{output_str}\033[0m" logging.info(output_str) @@ -2301,6 +2572,12 @@ def schedule_requests(self) -> int: nvtx_range_pop("add_request") elif header == Headers.SET_GENERATION_EPOCH: new_generation_epoch = data[1] + elif header == Headers.START_CUDA_PROFILER: + # Side-effect, not a state transition: apply immediately on every + # rank so an outer nsys --capture-range=cudaProfilerApi starts here. + torch.cuda.cudart().cudaProfilerStart() + elif header == Headers.STOP_CUDA_PROFILER: + torch.cuda.cudart().cudaProfilerStop() else: # Control signal: queue for second pass. self._pending_signals.append(message) diff --git a/megatron/core/inference/headers.py b/megatron/core/inference/headers.py index 8ad1913e6b1..4acef82bc1c 100644 --- a/megatron/core/inference/headers.py +++ b/megatron/core/inference/headers.py @@ -1,30 +1,137 @@ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. -from enum import Enum, auto +"""Message headers for inference coordinator/engine/client communication. +Headers are grouped into category IntEnum classes by concern. Each category +occupies a distinct numeric range so that wire values never collide across +categories. Adding a new class of headers means adding a new category enum in +its own range and listing it in HEADER_ENUMS; the existing categories are left +untouched. + +On the wire a header travels as its integer value. decode_header maps an +integer back to the originating category member. +""" + +from enum import IntEnum + + +class Connection(IntEnum): + """Client <-> coordinator handshake.""" + + CONNECT = 0 + CONNECT_ACK = 1 + + +class Request(IntEnum): + """Inference request submission and completion reply.""" + + SUBMIT_REQUEST = 20 + ENGINE_REPLY = 21 + + +class Control(IntEnum): + """Runtime control signals broadcast to engines.""" + + PAUSE = 40 + UNPAUSE = 41 + SUSPEND = 42 + RESUME = 43 + SET_GENERATION_EPOCH = 44 + STOP = 45 + START_CUDA_PROFILER = 46 + STOP_CUDA_PROFILER = 47 + + +class Lifecycle(IntEnum): + """Process lifecycle signals.""" + + DISCONNECT = 60 + SHUTDOWN = 61 + + +class Transport(IntEnum): + """Low-level transport framing.""" + + TP_BROADCAST = 80 -class Headers(Enum): - """ - Enum representing headers used for communication with the inference-coordinator. - """ - CONNECT = auto() - CONNECT_ACK = auto() - SUBMIT_REQUEST = auto() - ENGINE_REPLY = auto() - PAUSE = auto() - UNPAUSE = auto() - SUSPEND = auto() - RESUME = auto() - SET_GENERATION_EPOCH = auto() - STOP = auto() - DISCONNECT = auto() - SHUTDOWN = auto() - TP_BROADCAST = auto() +# All header categories. To add a new class of headers, define a new IntEnum in +# its own (disjoint) numeric range and append it here; nothing else needs to +# change. Ranges are spaced out to leave room for growth within each category. +HEADER_ENUMS = (Connection, Request, Control, Lifecycle, Transport) class UnknownHeaderError(Exception): - """A signal with an unrecognized header was received by the coordinator.""" + """A signal with an unrecognized header was received.""" def __init__(self, header): super().__init__(f"specialize for {header}.") + + +def _build_tables(): + """Index every header by wire value and by name, asserting no collisions.""" + by_value = {} + by_name = {} + for enum_cls in HEADER_ENUMS: + for member in enum_cls: + if member.value in by_value: + raise ValueError( + f"duplicate header wire value {member.value}: " + f"{by_value[member.value]!r} and {member!r}" + ) + if member.name in by_name: + raise ValueError( + f"duplicate header name {member.name!r}: " + f"{by_name[member.name]!r} and {member!r}" + ) + by_value[member.value] = member + by_name[member.name] = member + return by_value, by_name + + +_CODE_TO_MEMBER, _NAME_TO_MEMBER = _build_tables() + + +def decode_header(value): + """Resolve an integer wire value to its category header member. + + Args: + value (int): The integer header value read off the wire. + + Returns: + The corresponding category enum member (e.g. ``Control.PAUSE``). + + Raises: + UnknownHeaderError: if no header is registered for ``value``. + """ + try: + return _CODE_TO_MEMBER[value] + except KeyError: + raise UnknownHeaderError(value) + + +class _Headers: + """Flat, read-only union over every header category. + + Lets callers use a single name without caring which category a header lives + in: attribute access (``Headers.PAUSE``) resolves to the underlying category + member (``Control.PAUSE``), and calling it (``Headers(value)``) decodes a + wire value via decode_header. Because both return the canonical category + members, equality and dict-key lookups against the categories work + unchanged. + """ + + def __call__(self, value): + return decode_header(value) + + def __getattr__(self, name): + try: + return _NAME_TO_MEMBER[name] + except KeyError: + raise AttributeError(f"no header named {name!r}") + + def __iter__(self): + return iter(_NAME_TO_MEMBER.values()) + + +Headers = _Headers() diff --git a/megatron/core/inference/inference_client.py b/megatron/core/inference/inference_client.py index f5a78d1b16c..c3017737146 100644 --- a/megatron/core/inference/inference_client.py +++ b/megatron/core/inference/inference_client.py @@ -201,6 +201,19 @@ def unpause_engines(self) -> None: """Sends UNPAUSE to all engines. No synchronization needed.""" self._send_signal_to_engines(Headers.UNPAUSE) + def start_cuda_profiler(self) -> None: + """Sends START_CUDA_PROFILER to all engines via coordinator. + + Each engine calls ``torch.cuda.profiler.start()`` (cudaProfilerStart) on + its next loop iteration, so an outer ``nsys profile --capture-range= + cudaProfilerApi`` begins recording. No synchronization needed. + """ + self._send_signal_to_engines(Headers.START_CUDA_PROFILER) + + def stop_cuda_profiler(self) -> None: + """Sends STOP_CUDA_PROFILER to all engines (cudaProfilerStop).""" + self._send_signal_to_engines(Headers.STOP_CUDA_PROFILER) + def set_generation_epoch(self, generation_epoch: int): """Sends a signal to stamp all in-flight requests with the given generation epoch. diff --git a/megatron/core/inference/inference_request.py b/megatron/core/inference/inference_request.py index f9fd60f1033..d325a4e89a0 100644 --- a/megatron/core/inference/inference_request.py +++ b/megatron/core/inference/inference_request.py @@ -142,6 +142,10 @@ class InferenceRequest: sampling_params: Optional[SamplingParams] = None inference_parameters: Optional[SamplingParams] = None prompt_tokens: Optional[List[int]] = None + # Prompt token count. Always populated when serializing a finished request so the + # API can report usage.prompt_tokens even when the prompt_tokens tensor itself is + # dropped from the payload (see SamplingParams.return_prompt_tokens). + prompt_length: Optional[int] = None arrival_time: Optional[float] = None status: Optional[Status] = None encoder_prompt: Optional[str] = None @@ -380,6 +384,7 @@ class DynamicInferenceRequest(InferenceRequest): # Prefix caching fields block_size_tokens: Optional[int] = None # Block size for hash computation enable_prefix_caching: bool = False # Whether prefix caching is enabled + num_cached_tokens: int = 0 # Tokens served from prefix cache (set by context on first match) # Computed field - not passed by caller precomputed_block_hashes: List[int] = field(default_factory=list) @@ -440,13 +445,31 @@ def serialize(self): serialization. """ nvtx_range_push("DynamicInferenceRequest.serialize") + + # The prompt length is always reported (needed for usage.prompt_tokens), + # but the prompt_tokens tensor is dropped from the wire payload unless the + # client asked for it back (return_prompt_tokens). This keeps the large + # prompt tensor off the engine->coordinator->API path. Null it around + # super() so the tensor is never serialized, then restore local state. + prompt_len = len(self.prompt_tokens) if self.prompt_tokens is not None else None + drop_prompt = ( + self.prompt_tokens is not None + and self.sampling_params is not None + and not getattr(self.sampling_params, "return_prompt_tokens", False) + ) + saved_prompt_tokens = None + if drop_prompt: + saved_prompt_tokens = self.prompt_tokens + self.prompt_tokens = None + obj = super().serialize() obj["events"] = [e.serialize() for e in self.events] obj.pop("event_add_engine", None) + obj["prompt_length"] = prompt_len # Sanity check routing_indices: ndarray [total_tokens - 1, num_layers, topk] if self.routing_indices is not None: - total_tokens = len(self.prompt_tokens) + len(self.generated_tokens) + total_tokens = prompt_len + len(self.generated_tokens) # the last generated token does not undergo a forward pass # hence we expect routing indices for total_tokens - 1 assert self.routing_indices.shape[0] == total_tokens - 1, ( @@ -454,6 +477,9 @@ def serialize(self): f"total tokens {total_tokens-1}." ) + if drop_prompt: + self.prompt_tokens = saved_prompt_tokens + nvtx_range_pop("DynamicInferenceRequest.serialize") return obj @@ -740,6 +766,7 @@ def merge_lists(key): block_size_tokens=self.requests[0].block_size_tokens, enable_prefix_caching=self.requests[0].enable_prefix_caching, precomputed_block_hashes=self.requests[0].precomputed_block_hashes, + num_cached_tokens=self.requests[0].num_cached_tokens, ) return request diff --git a/megatron/core/inference/sampling/base.py b/megatron/core/inference/sampling/base.py index dceebb060a8..a011b7ab08d 100644 --- a/megatron/core/inference/sampling/base.py +++ b/megatron/core/inference/sampling/base.py @@ -21,8 +21,11 @@ def sample_kernel( n: int, context, *, + no_top_k: bool, + no_top_p: bool, gather_indices: Optional[Tensor] = None, token_to_request_index: Optional[Tensor] = None, + output: Optional[Tensor] = None, eager: bool = False, cache_key: Any = None, ) -> Tensor: @@ -32,13 +35,18 @@ def sample_kernel( logits: Logits tensor of shape `[>=n, vocab_size]`. n: Number of rows to sample. context: The active DynamicInferenceContext. + no_top_k, no_top_p: Required batch-level dispatch flags (whether NO active + request uses top-k / top-p). The caller computes them once from the + pinned CPU sampling metadata (see the controller's + `_active_requests_sampling_filter_flags`), so the kernel never has to. gather_indices: If provided, only sample from `logits[gather_indices[:n], :]`. token_to_request_index: Per-token request mapping; when set, sampling parameters are gathered per-token instead of per-request. - eager, cache_key: Consumed by `CudaGraphManager` when it wraps this kernel. + output: Optional caller-owned destination tensor of shape `[n]`. + eager, cache_key: Accepted for API symmetry; ignored (no CUDA graph). Returns: - Sampled token ids of shape `[n]`. Under CUDA graph replay, this is a static buffer. + Sampled token ids of shape `[n]`. """ ... @@ -57,12 +65,24 @@ def sample_speculative( """Sample tokens for the speculative-verify path. Decode requests contribute `1 + num_speculative_tokens` rows; prefill requests contribute 1. - Builds the per-token request mapping and dispatches to `sample_kernel`. - The `sample_kernel` is forced eager so its own `CudaGraphManager` wrapper does not fire. + Builds the per-token request mapping and dispatches to the return-valued `sample_kernel`. When `gather_indices` is supplied, the kernel selects via `logits[gather_indices[:n], :]`. When `gather_indices` is None, `required_logits` is expected to be already pre-gathered to the layout described above (e.g. when `materialize_only_last_token_logits=True` upstream). + + Args: + required_logits: Logits containing base and speculative rows. + num_decode: Number of decode requests. + num_prefill: Number of prefill requests. + num_speculative_tokens: Number of draft tokens per decode request. + context: The active DynamicInferenceContext. + gather_indices: Optional rows to gather from `required_logits`. + eager: Whether to bypass a wrapped CUDA graph. + cache_key: CUDA graph lookup key. + + Returns: + Sampled token IDs for all required base and speculative rows. """ # CudaGraphManager consumes these args, if it exists. del eager, cache_key @@ -80,10 +100,19 @@ def sample_speculative( torch.arange(num_decode, num_decode + num_prefill, device=device), ] ) + # Batch-level dispatch flags, required by `sample_kernel`. Read from the same + # pinned CPU sampling metadata as the controller's filter flags (sync-free): a + # filter is absent only when NO active request uses it. + active_request_count = context.total_request_count - context.paused_request_count + md = context.active_request_metadata + no_top_k = bool((md["top_k"][:active_request_count] == 0).all()) + no_top_p = bool((md["top_p"][:active_request_count] == 0.0).all()) return self.sample_kernel( required_logits, num_tokens, context, + no_top_k=no_top_k, + no_top_p=no_top_p, gather_indices=gather_indices, token_to_request_index=token_to_request_index, eager=True, @@ -91,13 +120,15 @@ def sample_speculative( @abstractmethod def log_probs_kernel( - self, logits: Tensor, temperature: Tensor, top_k: Tensor, top_p: Tensor + self, logits: Tensor, context, *, token_to_request_index: Optional[Tensor] = None ) -> Tensor: """Per-row log-probs of the distribution this backend samples from. Args: logits: `[num_rows, vocab_size]` raw logits. - temperature, top_k, top_p: `[num_rows]` per-row sampling params. + context: The active DynamicInferenceContext. + token_to_request_index: Optional per-row request mapping. When + omitted, each logits row maps to the request at the same index. Returns: `[num_rows, vocab_size]` log-probs; filtered-out tokens are `-inf`. diff --git a/megatron/core/inference/sampling/flashinfer_sampling.py b/megatron/core/inference/sampling/flashinfer_sampling.py index f7b85a8836e..95c28125751 100644 --- a/megatron/core/inference/sampling/flashinfer_sampling.py +++ b/megatron/core/inference/sampling/flashinfer_sampling.py @@ -11,32 +11,33 @@ flashinfer = None from megatron.core.inference.sampling.base import Sampling -from megatron.core.transformer.cuda_graphs import CudaGraphManager class FlashInferSampling(Sampling): - """Fused FlashInfer sampling, with optional CUDA graph capture/replay.""" + """FlashInfer sampling with per-step top-p-only / top-k-only / joint dispatch. + + Each step selects a kernel from the batch's active filters: the dedicated exact + top-p or top-k kernel when only one filter is in use, and the joint kernel only + for genuinely mixed batches. The dispatch flags are read from the pinned CPU + sampling metadata, so evaluating them costs no GPU sync. + + The sampler runs eagerly. Its kernel choice is data-dependent (it varies with + which filters the batch uses), so it cannot be captured in a CUDA graph; running + eagerly also lets the controller's seeded RNG generator advance its philox offset + normally between steps -- fresh randomness per step, reproducible from the seed. + (FlashInfer bakes the philox state into a graph as a by-value constant at capture, + so a captured sampler replays identical random numbers; see + https://www.linkedin.com/pulse/pinned-rng-drifting-crash-from-cuda-graph-chenyang-zhao-csuac/) + """ def __init__( self, vocab_size: int, rng: torch.Generator, config=None, enable_cuda_graph: bool = False ) -> None: + # `config` / `enable_cuda_graph` are accepted for factory API symmetry but + # intentionally unused: the sampler is never graphed (see class docstring). + del config, enable_cuda_graph self._vocab_size = vocab_size self._rng = rng - if enable_cuda_graph and config is not None and config.cuda_graph_impl == "local": - CudaGraphManager( - config, - self, - function_name="sample_kernel", - need_backward=False, - inline_capture=True, - ) - CudaGraphManager( - config, - self, - function_name="sample_speculative", - need_backward=False, - inline_capture=True, - ) def sample_kernel( self, @@ -44,31 +45,37 @@ def sample_kernel( n: int, context, *, + no_top_k: bool, + no_top_p: bool, gather_indices: Optional[Tensor] = None, token_to_request_index: Optional[Tensor] = None, + output: Optional[Tensor] = None, eager: bool = False, cache_key: Any = None, ) -> Tensor: - """FlashInfer fused top-k / top-p sampling kernel. + """Sample tokens, dispatching top-p-only / top-k-only / joint by filter flags. Args: logits: Logits tensor of shape `[>=n, vocab_size]`. n: Number of rows to sample. context: The active DynamicInferenceContext. + no_top_k, no_top_p: Required batch-level dispatch flags (whether NO active + request uses top-k / top-p). The caller computes them once from the + pinned CPU sampling metadata (the controller's + `_active_requests_sampling_filter_flags`). gather_indices: When set, sample from `logits[gather_indices[:n], :]`. - token_to_request_index: When set, sampling parameters are gathered per-token - rather than per-request (used by the speculative path). - eager, cache_key: Consumed by `CudaGraphManager` when it wraps this kernel. + token_to_request_index: When set, sampling parameters are gathered + per-token rather than per-request (speculative decoding path). + output: Optional caller-owned destination tensor of shape `[n]`. + eager, cache_key: Accepted for API symmetry; ignored (no CUDA graph). Returns: - Sampled token ids of shape `[n]`. Under CUDA graph replay, this is a static buffer. + Sampled token IDs in `output`, or a newly allocated tensor when it is not provided. """ - # CudaGraphManager consumes these args, if it exists. del eager, cache_key - # Read GPU sampling parameters from the per-step gpu_view mirror. The - # CPU source-of-truth (`active_request_metadata`) is pinned but resident - # on CPU, so reading it here would mix devices with `logits`. + # Per-row sampling params (GPU) for the kernel. gpu_view mirrors the pinned + # CPU `active_request_metadata` via the per-step coalesced H2D. gv = context.gpu_view if token_to_request_index is None: temperature = gv.temperature[:n] @@ -79,31 +86,85 @@ def sample_kernel( top_k = gv.top_k[token_to_request_index] top_p = gv.top_p[token_to_request_index] - # Clamp temperature to avoid division by 0. + # Temperature scale. `temperature` is a float32 tensor, so `bf16 logits / + # temperature` promotes `scaled` to fp32 -- the softmax / nucleus math must + # run in fp32 (a bf16 softmax over the vocab loses precision in exactly the + # tail region top-p depends on). The assert pins that guarantee. temperature = temperature.clamp(min=1e-6) if gather_indices is None: scaled = logits[:n] / temperature.unsqueeze(1) else: scaled = logits[gather_indices[:n], :] / temperature.unsqueeze(1) - probs = torch.softmax(scaled, dim=-1) - - # Sentinel values disable filtering: - # top_k=vocab_size keeps all tokens, top_p=1.0 keeps the full probability mass. - # TODO: Consider changing the disable flags in the `InferenceRequest`. - top_k_safe = top_k.masked_fill(top_k == 0, self._vocab_size) - top_p_safe = top_p.masked_fill(top_p == 0.0, 1.0) - output = torch.empty(n, device=logits.device, dtype=torch.int64) - output.copy_( - flashinfer.sampling.top_k_top_p_sampling_from_probs( - probs, top_k_safe, top_p_safe, generator=self._rng - ) - ) + assert scaled.dtype == torch.float32, f"sampling math must be fp32, got {scaled.dtype}" + + # `no_top_k` / `no_top_p` are the caller-supplied batch-level dispatch flags: + # a filter is absent only when NO active request uses it. Per-row sentinels + # disable a filter for a row (top_k=vocab keeps all tokens, top_p=1.0 keeps + # the full mass). Every kernel gets `self._rng` so sampling is seeded and its + # philox offset advances per launch. + if no_top_k and no_top_p: + # No nucleus / top-k filtering: sample the full temperature-scaled + # distribution. Use FlashInfer's kernel rather than torch.multinomial: + # multinomial forces a device-to-host sync, whereas sampling_from_probs + # stays on-device and keeps the RNG's philox offset advancing per launch. + probs = torch.softmax(scaled, dim=-1) + sampled_tokens = flashinfer.sampling.sampling_from_probs( + probs, deterministic=True, generator=self._rng + ).long() + elif no_top_k: + # Top-p only -> dedicated exact nucleus kernel. + probs = torch.softmax(scaled, dim=-1) + top_p_safe = top_p.masked_fill(top_p == 0.0, 1.0) + sampled_tokens = flashinfer.sampling.top_p_sampling_from_probs( + probs, top_p_safe, deterministic=True, generator=self._rng + ).long() + elif no_top_p: + # Top-k only -> dedicated exact top-k kernel. + probs = torch.softmax(scaled, dim=-1) + top_k_safe = top_k.masked_fill(top_k == 0, self._vocab_size) + sampled_tokens = flashinfer.sampling.top_k_sampling_from_probs( + probs, top_k_safe, deterministic=True, generator=self._rng + ).long() + else: + # Mixed batch (some top-k, some top-p, or requests using both) -> joint + # kernel, fed the temperature-scaled logits. + top_k_safe = top_k.masked_fill(top_k == 0, self._vocab_size) + top_p_safe = top_p.masked_fill(top_p == 0.0, 1.0) + sampled_tokens = flashinfer.sampling.top_k_top_p_sampling_from_logits( + scaled, top_k_safe, top_p_safe, deterministic=True, generator=self._rng + ).long() + + if output is None: + return sampled_tokens + output.copy_(sampled_tokens) return output def log_probs_kernel( - self, logits: Tensor, temperature: Tensor, top_k: Tensor, top_p: Tensor + self, logits: Tensor, context, *, token_to_request_index: Optional[Tensor] = None ) -> Tensor: - """Per-row log-probs of the FlashInfer top-k / top-p sampling distribution.""" + """Per-row log-probs of the FlashInfer top-k / top-p sampling distribution. + + Args: + logits (Tensor): Raw logits with shape `[num_rows, vocab_size]`. + context: Active dynamic inference context providing GPU sampling metadata. + token_to_request_index (Optional[Tensor]): Optional mapping from each + logits row to its request index. + + Returns: + Tensor: Per-row log probabilities for the processed distribution. + """ + gpu_view = context.gpu_view + if token_to_request_index is None: + num_rows = logits.size(0) + temperature = gpu_view.temperature[:num_rows] + top_k = gpu_view.top_k[:num_rows] + top_p = gpu_view.top_p[:num_rows] + else: + token_to_request_index = token_to_request_index.to(logits.device, non_blocking=True) + temperature = gpu_view.temperature[token_to_request_index] + top_k = gpu_view.top_k[token_to_request_index] + top_p = gpu_view.top_p[token_to_request_index] + temperature = temperature.clamp(min=1e-6) probs = torch.softmax(logits / temperature.unsqueeze(1), dim=-1) diff --git a/megatron/core/inference/sampling/torch_sampling.py b/megatron/core/inference/sampling/torch_sampling.py index f7f6f8cb662..e76d18059d8 100644 --- a/megatron/core/inference/sampling/torch_sampling.py +++ b/megatron/core/inference/sampling/torch_sampling.py @@ -114,14 +114,34 @@ def sample_from_logits( return sampled def log_probs_kernel( - self, logits: Tensor, temperature: Tensor, top_k: Tensor, top_p: Tensor + self, logits: Tensor, context, *, token_to_request_index: Optional[Tensor] = None ) -> Tensor: """Per-row log-probs of the temperature, top-k/top-p sampling distribution. Buckets rows by identical (temperature, top_k, top_p) and reuses `filter_logits` - (the same filter as `sample_from_logits`) so log-probs match how this backend - samples. `temperature`/`top_k`/`top_p` are per-row `[num_rows]` tensors. + (the same filter as `sample_from_logits`) so log-probs match how this backend samples. + + Args: + logits (Tensor): Raw logits with shape `[num_rows, vocab_size]`. + context: Active dynamic inference context providing CPU sampling metadata. + token_to_request_index (Optional[Tensor]): Optional CPU mapping from + each logits row to its request index. + + Returns: + Tensor: Per-row log probabilities for the processed distribution. """ + active_request_count = context.total_request_count - context.paused_request_count + metadata = context.active_request_metadata + if token_to_request_index is None: + temperature = metadata["temperature"][:active_request_count] + top_k = metadata["top_k"][:active_request_count] + top_p = metadata["top_p"][:active_request_count] + else: + assert not token_to_request_index.is_cuda + temperature = metadata["temperature"][:active_request_count][token_to_request_index] + top_k = metadata["top_k"][:active_request_count][token_to_request_index] + top_p = metadata["top_p"][:active_request_count][token_to_request_index] + temps = temperature.tolist() top_ks = top_k.tolist() top_ps = top_p.tolist() @@ -144,28 +164,34 @@ def sample_kernel( n: int, context, *, + no_top_k: bool, + no_top_p: bool, gather_indices: Optional[Tensor] = None, token_to_request_index: Optional[Tensor] = None, + output: Optional[Tensor] = None, eager: bool = False, cache_key: Any = None, ) -> Tensor: - """Bucket active requests by `(temperature, top_k, top_p)` and sample each bucket. + """Bucket active requests by sampling parameters and sample each bucket. Args: logits: Logits tensor of shape `[>=n, vocab_size]`. n: Number of rows to sample. context: The active DynamicInferenceContext. + no_top_k, no_top_p: Batch-level dispatch flags (part of the shared kernel + contract); ignored here since the exact per-bucket sort already handles + any top-k / top-p combination. gather_indices: When set, sample from `logits[gather_indices[:n], :]`. token_to_request_index: When set, the loop dispatches per-token rather than per-request (used by the speculative path). + output: Optional caller-owned destination tensor of shape `[n]`. eager: Accepted for API symmetry; ignored (TorchSampling has no graph wrapper). cache_key: Accepted for API symmetry; ignored. Returns: - Sampled token ids of shape `[n]`. + Sampled token IDs in `output`, or a newly allocated tensor when it is not provided. """ - # CudaGraphManager consumes these args, if it exists. - del eager, cache_key + del eager, cache_key, no_top_k, no_top_p # Group active requests into sampling buckets by (temperature, top_k, top_p). active_request_count = context.total_request_count - context.paused_request_count @@ -187,7 +213,8 @@ def sample_kernel( if gather_indices is not None: logits = logits[gather_indices[:n], :] - output = torch.empty(n, device=logits.device, dtype=torch.int64) + if output is None: + output = torch.empty(n, device=logits.device, dtype=torch.int64) token_list = [] indices_list = [] for idx_tensor, (_, temp, top_k, top_p) in zip(bucket_index_tensors, buckets): diff --git a/megatron/core/inference/sampling_params.py b/megatron/core/inference/sampling_params.py index 13bc8ac0d7b..f7f95060cef 100644 --- a/megatron/core/inference/sampling_params.py +++ b/megatron/core/inference/sampling_params.py @@ -34,6 +34,10 @@ class SamplingParams: None # List of strings that will stop generation when produced ) detokenize_stop_sequence: bool = False # Keep stop words and EOD in generated text + # Echo prompt token ids back in the response. When False (default), the engine + # drops prompt_tokens before serializing the finished request, saving the ZMQ + # transmission cost for long prompts. Opt in when the client needs them. + return_prompt_tokens: bool = False def __post_init__(self): """Ensure backward compatibility for return_prompt_top_n_logprobs. diff --git a/megatron/core/inference/text_generation_controllers/text_generation_controller.py b/megatron/core/inference/text_generation_controllers/text_generation_controller.py index b9eba10a5a4..cace7af6cc9 100644 --- a/megatron/core/inference/text_generation_controllers/text_generation_controller.py +++ b/megatron/core/inference/text_generation_controllers/text_generation_controller.py @@ -6,7 +6,7 @@ import functools from collections import defaultdict from dataclasses import dataclass -from typing import Any, Dict, List, Optional, OrderedDict, Tuple, Union +from typing import Any, Callable, Dict, List, Optional, OrderedDict, Tuple, Union import numpy as np import torch @@ -41,6 +41,7 @@ from megatron.core.transformer.moe.moe_layer import BaseMoELayer from megatron.core.transformer.moe.router_replay import RouterReplay, RouterReplayAction from megatron.core.transformer.moe.router_trace import get_moe_router_tracer +from megatron.core.transformer.moe.token_dispatcher_inference import NVLSAllGatherVDispatcher from megatron.core.transformer.utils import set_model_to_sequence_parallel from megatron.core.utils import ( accepts_parameter, @@ -72,26 +73,136 @@ @dataclass -class DecodeForwardPrimer: - """Track whether a decode forward is ready to sample.""" +class AsyncScheduleLogitsState: + """Track logits submitted for the next async-scheduling sample.""" - is_primed: bool = False + is_valid: bool = False cuda_graph_request_count: Optional[int] = None + token_row_indices: Optional[Tensor] = None - def mark_primed(self, cuda_graph_request_count: Optional[int]) -> None: - """Record that a decode forward has produced logits ready for sampling. + def set_pending( + self, cuda_graph_request_count: Optional[int], token_row_indices: Optional[Tensor] = None + ) -> None: + """Record logits submitted for the next sample. Args: cuda_graph_request_count (Optional[int]): CUDA graph request count - for the primed forward, or `None` when CUDA graphs were not used. + for the pending logits, or `None` when CUDA graphs were not used. + token_row_indices (Optional[Tensor]): Original GPU input row for each + logical token row in the pending forward. """ - self.is_primed = True + self.is_valid = True self.cuda_graph_request_count = cuda_graph_request_count + self.token_row_indices = token_row_indices def clear(self) -> None: - """Clear any primed-forward state.""" - self.is_primed = False + """Clear the pending logits state.""" + self.is_valid = False self.cuda_graph_request_count = None + self.token_row_indices = None + + +@dataclass +class _AsyncScheduleSampleResult: + """GPU samples, reusable CPU views, and readiness events for one async step.""" + + sampled_tokens_gpu: Tensor + sampled_tokens_cpu_view: Tensor + sampled_mtp_tokens_gpu: Optional[Tensor] + sampled_mtp_tokens_cpu_view: Optional[Tensor] + accepted_tokens_cpu_view: Optional[Tensor] + accepted_counts_gpu: Optional[Tensor] + accepted_counts_cpu_view: Optional[Tensor] + accepted_counts_cpu_ready_event: Optional[torch.cuda.Event] + sample_cpu_ready_event: Optional[torch.cuda.Event] + + +@dataclass(frozen=True) +class DecodeOnly: + """Decode-only state for the consumed and launched forwards. + + Attributes: + consumed: Whether the consumed output came from a decode-only forward, + or ``None`` when no output was consumed. + launched: Whether the launched forward is decode-only, or ``None`` when + no real forward was launched. + """ + + consumed: Optional[bool] + launched: Optional[bool] + + def __bool__(self) -> bool: + """Return the shared decode-only state when both forwards agree. + + Returns: + bool: The common consumed and launched decode-only state. + + Raises: + ValueError: If either forward is absent or the two states differ. + """ + if self.consumed is None or self.launched is None or self.consumed != self.launched: + raise ValueError( + "Decode-only state is ambiguous: " + f"consumed={self.consumed}, launched={self.launched}." + ) + return self.consumed + + +@dataclass(frozen=True) +class DynamicBatchControllerStepResult: + """Result of one dynamic-batching controller step. + + Attributes: + decode_only: Decode-only state for the consumed and launched forwards. + output: Sampled-step output, or ``None`` when no output was produced. + primer_only: Whether the step launched only an async-scheduling primer. + """ + + decode_only: DecodeOnly + output: Optional[Dict] = None + primer_only: bool = False + + +@dataclass +class _AsyncScheduleRequestResult: + """Request state produced by async scheduling bookkeeping.""" + + sampled_tokens_cpu: Tensor + accepted_tokens_cpu: Optional[Tensor] + active_request_ids: Tensor + finished_request_ids: Tensor + survivor_idxs: Optional[Tensor] = None + newly_paused_request_ids: Optional[Tensor] = None + evict_request_ids: Optional[Tensor] = None + + +@dataclass +class _AsyncScheduleLogProbsGPUResult: + """GPU logprob outputs awaiting transfer to CPU.""" + + selected_log_probs: Tensor + top_n_log_probs: Optional[Tensor] + top_n_token_ids: Optional[Tensor] + row_counts: List[int] + top_n_counts: List[int] + skip_prompt_log_probs: List[bool] + num_decode_requests: int + gpu_ready_event: Optional[torch.cuda.Event] + + +@dataclass +class _AsyncScheduleLogProbsTransfer: + """Transient CPU views retaining their GPU sources until D2H completes.""" + + selected_log_probs_cpu_view: Tensor + top_n_log_probs_cpu_view: Optional[Tensor] + top_n_token_ids_cpu_view: Optional[Tensor] + row_counts: List[int] + top_n_counts: List[int] + skip_prompt_log_probs: List[bool] + num_decode_requests: int + cpu_ready_event: Optional[torch.cuda.Event] + gpu_result: _AsyncScheduleLogProbsGPUResult # pylint: disable=line-too-long @@ -116,8 +227,10 @@ def __init__(self, inference_wrapped_model: AbstractModelInferenceWrapper, token pg_collection = inference_config.pg_collection if pg_collection is not None: self.pp_group = pg_collection.pp + self.dp_group = pg_collection.dp else: self.pp_group = parallel_state.get_pipeline_model_parallel_group() + self.dp_group = parallel_state.get_data_parallel_group() self.model_is_pipeline_parallel = self.model_config.pipeline_model_parallel_size > 1 @@ -129,8 +242,20 @@ def __init__(self, inference_wrapped_model: AbstractModelInferenceWrapper, token else: self.vocab_size = unwrapped_model.vocab_size + # Build and seed sampling RNG. Optionally offset by DP rank so each rank gets a + # unique generation seed (avoids identical samples when the same prompt is + # assigned to multiple DP ranks, which can corrupt RL training). Controlled by + # InferenceConfig.offset_sampling_seed_by_dp_rank, but deactivated when enabling + # --deterministic-mode (model_config.deterministic_mode). self.sampling_rng = torch.Generator(device=torch.cuda.current_device()) - self.sampling_rng.manual_seed(self.model_config.inference_sampling_seed) + seed = self.model_config.inference_sampling_seed + offset_by_dp = ( + inference_config.offset_sampling_seed_by_dp_rank + and not self.model_config.deterministic_mode + ) + if offset_by_dp: + seed += torch.distributed.get_rank(group=self.dp_group) + self.sampling_rng.manual_seed(seed) if not self.num_speculative_tokens: self.num_mtp_depths = 0 @@ -195,14 +320,24 @@ def _init_dynamic_sampling_tensors(self): ) else: self._all_logits_cuda = None - self._decode_forward_primer = DecodeForwardPrimer() - # Speculative path: - # - `self._sampled_tokens_cuda` is pre-allocated by `_init_mtp_sampling_tensors`. - # - The tensor cannot be reused between the Triton kernel and the sampling graph. - # Non-speculative path: - # - `self._sampled_tokens_cuda` is rebound to the output of `sample_kernel`, - # which uses CudaGraphManager syntactic sugar to keep it as a static tensor. - self._sampled_tokens_cuda = None + self._async_sched_logits = AsyncScheduleLogitsState() + # This buffer has a stable address across legacy, no-overlap, overlap, + # and MTP routing. Sampling producers must copy into it rather than rebind it. + self._sampled_tokens_cuda = torch.empty(max_requests, dtype=torch.int64, device=device) + self._async_sched_sampled_tokens_cpu_buffer = torch.empty( + max_requests, dtype=torch.int64, device="cpu", pin_memory=True + ) + self._async_sched_selected_log_probs_cpu_buffer = torch.empty( + context.max_tokens, dtype=torch.float32, device="cpu", pin_memory=True + ) + self._async_sched_top_n_log_probs_cpu_buffer = None + self._async_sched_top_n_token_ids_cpu_buffer = None + self._async_sched_top_n_capacity = 0 + self._async_sched_sample_gpu_ready_event = torch.cuda.Event() + self._async_sched_sample_cpu_ready_event = torch.cuda.Event() + self._async_sched_log_probs_gpu_ready_event = torch.cuda.Event() + self._async_sched_log_probs_cpu_ready_event = torch.cuda.Event() + self._async_sched_copy_stream = torch.cuda.Stream(device=device) # Sampling backend: provides the sampling kernel. if self._sampling_backend == "flashinfer": @@ -228,19 +363,26 @@ def _init_mtp_sampling_tensors(self): Addresses must be stable across steps for CUDA graph capture. """ + self._mtp_resolved_padded_count = None if not self.num_speculative_tokens: self._sampled_mtp_tokens_cuda = None self._accepted_tokens_per_request = None self._last_accepted_seq_indices = None + self._async_sched_mtp_token_row_indices = None + self._async_sched_sampled_mtp_tokens_cpu_buffer = None + self._async_sched_accepted_tokens_cpu_buffer = None + self._async_sched_accepted_counts_cpu_buffer = None + self._async_sched_mtp_verification_gpu_ready_event = None + self._async_sched_accepted_counts_cpu_ready_event = None return context = self.inference_wrapped_model.inference_context max_requests = context.max_requests device = torch.cuda.current_device() - self._sampled_tokens_cuda = torch.empty(max_requests, dtype=torch.int64, device=device) self._sampled_mtp_tokens_cuda = torch.empty( [self.num_speculative_tokens, max_requests], dtype=torch.int64, device=device ) + self._async_sched_mtp_token_row_indices = torch.arange(context.max_tokens, device=device) self._accepted_tokens_per_request = ( torch.ones( [max_requests, self.num_speculative_tokens], dtype=torch.int64, device=device @@ -258,69 +400,23 @@ def _init_mtp_sampling_tensors(self): self._mtp_position_ids_buf = torch.empty( [1, max_requests], dtype=torch.int64, device=device ) - - def _validate_async_sched_support_for_step(self) -> None: - """Validate controller/context state for async scheduling. - - Raises if the current step does not support async scheduling. - """ - context = self.inference_wrapped_model.inference_context - if not context.config.materialize_only_last_token_logits: - raise RuntimeError("Async scheduling requires materialize_only_last_token_logits=True.") - if self.num_speculative_tokens != 0: - raise RuntimeError("Async scheduling does not support speculative tokens.") - if context.is_hybrid_model: - raise RuntimeError("Async scheduling does not support hybrid/Mamba models.") - if context.enable_prefix_caching: - raise RuntimeError("Async scheduling does not support prefix caching.") - if context.paused_request_count != 0: - raise RuntimeError("Async scheduling does not support paused requests.") - if context.chunked_prefill_request_id != -1: - raise RuntimeError("Async scheduling does not support chunked prefill.") - if self.model_config.expert_model_parallel_size > 1: - raise RuntimeError("Async scheduling does not support expert parallelism.") - if self.model_config.num_moe_experts is not None: - raise RuntimeError("Async scheduling does not support MoE models.") - if self.model_config.moe_enable_routing_replay: - raise RuntimeError("Async scheduling does not support routing replay.") - - active_request_count = context.total_request_count - context.paused_request_count - active_slice = slice(context.paused_request_count, context.total_request_count) - if active_request_count == 0: - return - if not torch.all(context.request_metadata["top_k"][active_slice] == 1): - raise RuntimeError( - "Async scheduling only supports greedy sampling " "(SamplingParams.top_k == 1)." - ) - if not torch.all(context.request_metadata["top_p"][active_slice] == 0.0): - raise RuntimeError( - "Async scheduling only supports greedy sampling " "(SamplingParams.top_p == 0.0)." - ) - if torch.any(context.request_metadata["return_log_probs"][active_slice]): - raise RuntimeError("Async scheduling does not support log probabilities.") - if torch.any(context.request_metadata["top_n_logprobs"][active_slice] > 0): - raise RuntimeError("Async scheduling does not support top-n log probabilities.") - - def _compact_async_sched_logits(self, survivor_idxs: Tensor) -> None: - """Compact cached logits from old active-row order into survivor order. - - Args: - survivor_idxs (Tensor): Active-row indices for requests that remain - active after async scheduling. - """ - if survivor_idxs.numel() == 0: - self._decode_forward_primer.clear() - return - - survivor_idxs_cuda = survivor_idxs.to(self._all_logits_cuda.device) - compacted_logits = self._all_logits_cuda[:, survivor_idxs_cuda, :].contiguous() - if self._enable_cuda_graph: - self._all_logits_cuda[:, : survivor_idxs.numel(), :].copy_(compacted_logits) - else: - self._all_logits_cuda = compacted_logits - self._decode_forward_primer.mark_primed( - self._decode_forward_primer.cuda_graph_request_count + self._async_sched_sampled_mtp_tokens_cpu_buffer = torch.empty( + [self.num_speculative_tokens, max_requests], + dtype=torch.int64, + device="cpu", + pin_memory=True, + ) + self._async_sched_accepted_tokens_cpu_buffer = torch.empty( + [max_requests, self.num_speculative_tokens], + dtype=torch.int64, + device="cpu", + pin_memory=True, + ) + self._async_sched_accepted_counts_cpu_buffer = torch.empty( + max_requests, dtype=torch.int64, device="cpu", pin_memory=True ) + self._async_sched_mtp_verification_gpu_ready_event = torch.cuda.Event() + self._async_sched_accepted_counts_cpu_ready_event = torch.cuda.Event() @staticmethod def tokenize_prompt(tokenizer, prompt: str, add_BOS: bool = False) -> List[int]: @@ -621,17 +717,24 @@ def _dynamic_step_context_init( self, construct_graph_dimensions: Optional[InferenceBatchDimensions] = None, is_dummy_forward: bool = False, - ): + transfer_bookkeeping_to_gpu: bool = True, + record_bookkeeping_done_event: bool = False, + ) -> Tuple[Tensor, Tensor, Optional[torch.cuda.Event]]: """Initializes the inference context for dynamic batching. Args: construct_graph_dimensions (Optional[InferenceBatchDimensions]): The graph config to use for constructing the cuda graphs. is_dummy_forward (bool): Whether we are running an expert parallel dummy forward pass + transfer_bookkeeping_to_gpu (bool): Whether to publish the prepared + CPU bookkeeping snapshot to GPU before returning. + record_bookkeeping_done_event (bool): Whether to record an event + after the bookkeeping H2D transfer. - Return: - input_ids (Tensor): The active input IDs. - position_ids (Tensor): The active position IDs. + Returns: + Tuple[Tensor, Tensor, Optional[torch.cuda.Event]]: The active input + IDs, position IDs, and optional bookkeeping H2D completion + event. """ context = self.inference_wrapped_model.inference_context @@ -639,19 +742,16 @@ def _dynamic_step_context_init( unwrapped_model = unwrap_model(self.inference_wrapped_model.model) model_config = get_model_config(unwrapped_model) - # Initialize attention state (100% CPU computation). + # Initialize attention state and optionally publish CPU bookkeeping to GPU. range_push("initialize_attention_state") - context.initialize_attention_state( + bookkeeping_done_event = context.initialize_attention_state( construct_graph_dimensions=construct_graph_dimensions, is_expert_parallel_dummy_cuda_graph_step=is_dummy_forward, + transfer_bookkeeping_to_gpu=transfer_bookkeeping_to_gpu, + record_bookkeeping_done_event=record_bookkeeping_done_event, ) range_pop() - # Single batch CPU-to-GPU transfer of bookkeeping state. - range_push("transfer_bookkeeping_to_gpu") - context.transfer_bookkeeping_to_gpu() - range_pop() - set_moe_metadata_sync(unwrapped_model) # Derive the MTP padded batch size from the existing padded graph dimensions. @@ -699,11 +799,12 @@ def _dynamic_step_context_init( # If we are running a dummy forward step we want to use the token count agreed upon # by all EP ranks rather than the minimum number of tokens. if construct_graph_dimensions is not None and not is_dummy_forward: - return context.current_input_and_position_ids( + input_ids, position_ids = context.current_input_and_position_ids( num_warmup_tokens=construct_graph_dimensions.token_count ) else: - return context.current_input_and_position_ids() + input_ids, position_ids = context.current_input_and_position_ids() + return input_ids, position_ids, bookkeeping_done_event def _dynamic_step_forward_logits(self, input_ids: Tensor, position_ids: Tensor): """Forward step the model to get logits for dynamic batching. @@ -757,43 +858,7 @@ def _dynamic_step_forward_logits(self, input_ids: Tensor, position_ids: Tensor): else: self._all_logits_cuda = logits - def _run_async_sched_prepare(self, new_sample_copy: Tensor) -> Tuple[Tensor, Tensor]: - """Prepare decode requests and GPU-visible forward state for async scheduling. - - Args: - new_sample_copy (Tensor): CPU copy of sampled tokens for active requests. - - Returns: - Tuple[Tensor, Tensor]: Input token IDs and position IDs for the speculative forward. - """ - context = self.inference_wrapped_model.inference_context - context.prepare_requests(new_sample_copy) - return self._dynamic_step_context_init() - - def _run_async_sched_forward(self, input_ids: Tensor, position_ids: Tensor) -> Optional[int]: - """Run one dynamic forward pass and cache logits for async scheduling. - - Args: - input_ids (Tensor): The input token IDs. - position_ids (Tensor): The position IDs. - - Returns: - Optional[int]: CUDA graph request count for the forward pass, or - `None` when CUDA graphs were not used. - """ - context = self.inference_wrapped_model.inference_context - cuda_graph_request_count = ( - context.padded_active_request_count if context.using_cuda_graph_this_step() else None - ) - - range_push("forward_pass") - self._dynamic_step_forward_logits(input_ids, position_ids) - range_pop() - - self._decode_forward_primer.mark_primed(cuda_graph_request_count) - return cuda_graph_request_count - - def _rewind_kv_cache(self) -> tuple: + def _rewind_kv_cache(self, accepted_counts_cpu: Optional[Tensor] = None) -> tuple: """Update the KV cache bookkeeping for speculative decoding. After forward pass with speculative tokens, some tokens may be rejected. @@ -802,8 +867,12 @@ def _rewind_kv_cache(self) -> tuple: CPU source-of-truth tensors in place); the Mamba hybrid-model state update stays on GPU because it operates on GPU-resident state buffers. - Returns (blocks_to_release, remove_mask) for the caller to release blocks - back to the allocator outside the compiled graph. + Args: + accepted_counts_cpu (Optional[Tensor]): Accepted MTP draft counts already + copied to CPU. When omitted, this method performs the legacy D2H copy. + + Returns: + tuple: Blocks detached by rewind and the mask selecting valid block IDs. """ context = self.inference_wrapped_model.inference_context active_request_count = context.total_request_count - context.paused_request_count @@ -811,12 +880,13 @@ def _rewind_kv_cache(self) -> tuple: # accepted_counts is the only GPU input; D2H a small slice so the # CPU rewind can read its values via .tolist() inside a Python loop. - accepted_tokens_per_request_cpu = self._accepted_token_counts_per_request[ - :active_request_count - ].cpu() + if accepted_counts_cpu is None: + accepted_counts_cpu = self._accepted_token_counts_per_request[ + :active_request_count + ].cpu() blocks_to_release, remove_mask = rewind_kv_cache( - accepted_counts=accepted_tokens_per_request_cpu, + accepted_counts=accepted_counts_cpu, prefill_status=context.request_in_prefill_status_tensor[active_request_slice], last_kv_block_offset=context.request_last_kv_block_offset[active_request_slice], kv_length_offsets=context.request_kv_length_offsets[active_request_slice], @@ -868,14 +938,17 @@ def _sample_from_logits_2d(self, logits_2d: Tensor) -> Tensor: Returns: Tensor: Sampled tokens of shape [num_requests]. """ + no_top_k, no_top_p = self._active_requests_sampling_filter_flags() return self._sampling.sample_kernel( logits_2d, logits_2d.shape[0], self.inference_wrapped_model.inference_context, + no_top_k=no_top_k, + no_top_p=no_top_p, eager=True, ) - def _compute_serial_mtp_and_sample(self): + def _compute_serial_mtp_and_sample(self, base_position: Optional[Tensor] = None) -> None: """Compute MTP logits serially after verification and sample speculative tokens. This ensures that MTP predictions are always conditioned on verified tokens. @@ -886,6 +959,10 @@ def _compute_serial_mtp_and_sample(self): When sequence parallelism is active, hidden states are kept in SP format (scattered along the first dimension) between MTP depths to avoid a redundant gather + scatter round-trip per depth. + + Args: + base_position (Optional[Tensor]): GPU position of the first new MTP draft + for each request. Legacy scheduling derives it from rewound CPU state. """ nvtx_range_push("mtp-spec-decoding/serial-mtp-init") context = self.inference_wrapped_model.inference_context @@ -915,26 +992,23 @@ def _compute_serial_mtp_and_sample(self): else: last_accepted_hidden = None - # Compute position IDs for the next tokens. - # After rewind, request_kv_length_offsets has been adjusted. Read from - # CPU context (post-rewind values), NOT gpu_view (stale pre-rewind snapshot). - # The next position to predict is: adjusted_offset + processed_tokens. - cuda_device = torch.cuda.current_device() - adjusted_offsets = context.request_kv_length_offsets[active_slice].to( - cuda_device, non_blocking=True - ) - processed_tokens = context.request_query_lengths[active_slice].to( - cuda_device, non_blocking=True - ) - # Cast to int64 to match CUDA graph capture dtype expectations. - base_position = (adjusted_offsets + processed_tokens).to(torch.int64) + if base_position is None: + # Legacy scheduling derives positions from post-rewind CPU state. + cuda_device = torch.cuda.current_device() + adjusted_offsets = context.request_kv_length_offsets[active_slice].to( + cuda_device, non_blocking=True + ) + processed_tokens = context.request_query_lengths[active_slice].to( + cuda_device, non_blocking=True + ) + base_position = (adjusted_offsets + processed_tokens).to(torch.int64) # Start with the freshly sampled base token. next_token_ids = self._sampled_tokens_cuda[:active_request_count].clone() current_hidden = last_accepted_hidden if has_mtp else None # Compute padding needed to make batch compatible with SP and CUDA graphs. - if getattr(self, '_mtp_resolved_padded_count', None) is not None: + if self._mtp_resolved_padded_count is not None: # CUDA-graph path: use the EP-synced padded count. padded_count = self._mtp_resolved_padded_count assert not self._sp_enabled or padded_count % self._tp_size == 0 @@ -963,6 +1037,14 @@ def _compute_serial_mtp_and_sample(self): position_ids_buf[0, active_request_count:] = 0 nvtx_range_pop("mtp-spec-decoding/serial-mtp-init") + + # MTP MoE forwards are request-count shaped: the routing map holds + # active_request_count real rows followed by padding up to padded_count. + # The NVLS routing mask defaults to the main step's token count, so point + # it at the MTP row count instead, else padding rows route to experts. + if context._nvls_dispatcher: + NVLSAllGatherVDispatcher.modify_real_token_count_for_mtp(active_request_count) + for depth in range(self.num_mtp_depths): nvtx_range_push(f"mtp-spec-decoding/depth-{depth}") @@ -1036,31 +1118,30 @@ def _verify_speculative_tokens( num_speculative_tokens=self.num_speculative_tokens, ) - def _dynamic_step_sample_logits_and_verify_tokens(self, input_ids: Tensor): - """ - Sample tokens from logits for dynamic batching with speculative tokens and verify the tokens. + def _dynamic_step_sample_logits_and_verify_tokens( + self, input_ids: Tensor, token_row_indices: Optional[Tensor] = None + ) -> None: + """Sample MTP logits and verify pending draft tokens. + + Args: + input_ids (Tensor): Input token storage used by the pending forward. + token_row_indices (Optional[Tensor]): Original GPU input row for each + current logical token row after survivor compaction. """ context = self.inference_wrapped_model.inference_context active_request_count = context.total_request_count - context.paused_request_count - # Sampling-side request counts: padded when running a captured graph. - # Verify uses the actual counts so the Triton kernels operate on the real workload. - use_graph_for_sampling = ( - self._sampling_backend == "flashinfer" - and self._enable_cuda_graph - and context.using_cuda_graph_this_step() - ) - if use_graph_for_sampling: - sample_num_decode = context.padded_batch_dimensions.decode_req_count - sample_num_prefill = context.padded_batch_dimensions.prefill_req_count - else: - sample_num_decode = context.num_decode_requests - sample_num_prefill = context.num_prefill_requests + # The FlashInfer sampler runs eagerly (never CUDA-graphed), so verify with the + # actual request counts. When the forward pass is graphed `required_logits` is + # padded to a static shape, but sampling only the actual token prefix (below) + # leaves the trailing padded rows unsampled. + sample_num_decode = context.num_decode_requests + sample_num_prefill = context.num_prefill_requests # Logit indices for tokens that need sampling. - # Padded under graph capture so the captured `gather_indices` input has a stable shape. - # Padded slots resolve to row 0; verify and prepare-next read only the actual prefix, - # so the padded-row samples produced by the captured kernel are discarded. + # `speculative_required_logit_indices()` pads to a static shape when the forward + # pass is graphed (trailing slots resolve to row 0); sampling uses the actual + # counts, so those padded slots are never sampled. nvtx_range_push("mtp-spec-decoding/verify/logit-indices") # Use pre-allocated buffer for CUDA graph compatibility. logits = self._all_logits_cuda @@ -1090,12 +1171,6 @@ def _dynamic_step_sample_logits_and_verify_tokens(self, input_ids: Tensor): self.num_speculative_tokens, context, gather_indices=sample_gather_indices, - eager=not use_graph_for_sampling, - cache_key=( - ("sample_speculative", sample_num_decode, sample_num_prefill) - if use_graph_for_sampling - else None - ), ) nvtx_range_pop("mtp-spec-decoding/verify/sample") @@ -1104,7 +1179,13 @@ def _dynamic_step_sample_logits_and_verify_tokens(self, input_ids: Tensor): # Verify speculative tokens against input tokens. nvtx_range_push("mtp-spec-decoding/verify/verify-tokens") - input_tokens_required = input_ids[0, required_logit_indices] + input_row_indices = required_logit_indices + if token_row_indices is not None: + actual_required_count = active_request_count * (self.num_speculative_tokens + 1) + input_row_indices = token_row_indices[ + required_logit_indices[:actual_required_count].long() + ] + input_tokens_required = input_ids[0, input_row_indices] last_one_indices, accepted_tokens_mask, input_tokens_required = ( self._verify_speculative_tokens( output_tokens, @@ -1120,7 +1201,7 @@ def _dynamic_step_sample_logits_and_verify_tokens(self, input_ids: Tensor): self._prepare_speculative_tokens_for_next_forward_pass( num_decode_requests, output_tokens, - required_logit_indices, + input_row_indices, last_one_indices, accepted_tokens_mask, input_tokens_required, @@ -1168,13 +1249,9 @@ def _dynamic_step_sample_logits(self): context = self.inference_wrapped_model.inference_context active_request_count = context.total_request_count - context.paused_request_count - use_graph = ( - self._sampling_backend == "flashinfer" - and self._enable_cuda_graph - and context.using_cuda_graph_this_step() - ) - # Padded count when running a captured graph (cache key buckets); actual otherwise. - n = context.padded_active_request_count if use_graph else active_request_count + # The FlashInfer sampler runs eagerly (never CUDA-graphed), so sample the + # actual active rows -- there is no captured static shape to pad up to. + n = active_request_count # When `materialize_only_last_token_logits` is true the forward pass already # selected the right rows. Otherwise we point the kernel at the per-request # last-token positions via `gather_indices`; padded slots safely fan in to row 0. @@ -1183,27 +1260,57 @@ def _dynamic_step_sample_logits(self): if context.config.materialize_only_last_token_logits else context.gpu_view.active_request_last_token_idxs ) - self._sampled_tokens_cuda = self._sampling.sample_kernel( + no_top_k, no_top_p = self._active_requests_sampling_filter_flags(active_request_count) + self._sampling.sample_kernel( self._all_logits_cuda.squeeze(0), n, context, gather_indices=gather_indices, - eager=not use_graph, - cache_key=("sample", n) if use_graph else None, + no_top_k=no_top_k, + no_top_p=no_top_p, + output=self._sampled_tokens_cuda[:n], + ) + + def _active_requests_sampling_filter_flags( + self, active_request_count: Optional[int] = None + ) -> Tuple[bool, bool]: + """Return ``(no_top_k, no_top_p)`` batch-level escape hatches for the active batch. + + These drive the FlashInfer sampler's dispatch (top-p-only / top-k-only / + joint) and are read from the pinned CPU sampling metadata, so they incur no + GPU sync. A filter is "absent" only when NO active request uses it. Padded + rows carry a neutral 0 and never flip a flag. + """ + context = self.inference_wrapped_model.inference_context + active_request_count = ( + context.total_request_count - context.paused_request_count + if active_request_count is None + else active_request_count ) + if active_request_count <= 0: + return True, True + + active_metadata = context.active_request_metadata + active_slice = slice(0, active_request_count) + no_top_k = bool((active_metadata["top_k"][active_slice] == 0).all()) + no_top_p = bool((active_metadata["top_p"][active_slice] == 0.0).all()) + return no_top_k, no_top_p def _dynamic_step_log_probs_bookkeeping(self) -> Tuple[bool, bool]: """Perform bookkeeping necessary to compute log probs for dynamic batching. Returns: return_log_probs (bool): Whether to return the sampled log_probs. + return_top_n_logprobs (bool): Whether to return top-n log_probs. """ context = self.inference_wrapped_model.inference_context active_request_count = context.total_request_count - context.paused_request_count return ( - (context.active_request_metadata["return_log_probs"][:active_request_count]).any(), - (context.active_request_metadata["top_n_logprobs"][:active_request_count] > 0).any(), + bool(context.active_request_metadata["return_log_probs"][:active_request_count].any()), + bool( + (context.active_request_metadata["top_n_logprobs"][:active_request_count] > 0).any() + ), ) def _router_record_bookkeeping(self) -> Optional[np.ndarray]: @@ -1596,36 +1703,17 @@ def _dynamic_step_calculate_top_n_logprobs( return top_n_results if top_n_results else None - @torch.inference_mode() - def dummy_forward(self): - """Perform a dummy forward pass. This is used in expert model parallelism - on ranks that do not have any real requests. It may run in eager mode.""" - - context = self.inference_wrapped_model.inference_context + def _run_dummy_base_forward(self, input_ids: Tensor, position_ids: Tensor) -> None: + """Run the base-model portion of an expert-parallel dummy step. - # attempt to use cuda-graph if possible - input_ids, position_ids = self._dynamic_step_context_init(is_dummy_forward=True) + Args: + input_ids (Tensor): Dummy input token IDs. + position_ids (Tensor): Dummy input position IDs. + """ self._dynamic_step_forward_logits(input_ids, position_ids) - # Disable MoE padding for MTP computation, unless CUDA graphs - # are active (the graphs were captured with padding enabled). - if self.model_config.moe_pad_experts_for_cuda_graph_inference: - if not context.using_cuda_graph_this_step(): - unwrapped_model = unwrap_model(self.inference_wrapped_model.model) - set_decode_expert_padding(unwrapped_model, False) - - # When speculative decoding is active, the real EP ranks perform serial - # MTP forward passes after the main forward pass. MTP layers may contain - # MoE sublayers (inherited from the decoder spec), which require EP - # all-to-all collectives. The dummy rank must participate in these - # collectives to avoid a hang. - self._dummy_serial_mtp_forward() - - # clear the context of any temporary state from the dummy forward - context.reset() - @torch.inference_mode() - def _dummy_serial_mtp_forward(self): + def _run_dummy_serial_mtp_forward(self) -> None: """Run dummy MTP forward passes to participate in EP collectives. When speculative decoding is active and MTP layers contain MoE sublayers @@ -1644,11 +1732,9 @@ def _dummy_serial_mtp_forward(self): if self.model_config.expert_model_parallel_size <= 1: return + context = self.inference_wrapped_model.inference_context unwrapped_model = self._unwrapped_model - - has_mtp = self._is_last_pp_stage and hasattr( - unwrapped_model, '_decoder_hidden_states_cache' - ) + has_mtp = self._is_last_pp_stage and hasattr(unwrapped_model, "mtp") if not has_mtp and not self.model_is_pipeline_parallel: # No MTP on this rank and no PP broadcast to participate in. return @@ -1659,7 +1745,7 @@ def _dummy_serial_mtp_forward(self): # Use precomputed MTP CUDA graph batch size when available; # otherwise use minimal SP-compatible size. - if getattr(self, '_mtp_resolved_padded_count', None) is not None: + if self._mtp_resolved_padded_count is not None: padded_count = self._mtp_resolved_padded_count assert not self._sp_enabled or padded_count % self._tp_size == 0 elif has_mtp: @@ -1709,6 +1795,58 @@ def _dummy_serial_mtp_forward(self): ) nvtx_range_pop(f"mtp-spec-decoding/dummy-depth-{depth}") + def _run_dummy_legacy_step(self, input_ids: Tensor, position_ids: Tensor) -> None: + """Run a legacy dummy step in base-forward then MTP order. + + Args: + input_ids (Tensor): Dummy input token IDs. + position_ids (Tensor): Dummy input position IDs. + """ + context = self.inference_wrapped_model.inference_context + self._run_dummy_base_forward(input_ids, position_ids) + + # Disable MoE padding for MTP computation, unless CUDA graphs + # are active (the graphs were captured with padding enabled). + if self.model_config.moe_pad_experts_for_cuda_graph_inference: + if not context.using_cuda_graph_this_step(): + unwrapped_model = unwrap_model(self.inference_wrapped_model.model) + set_decode_expert_padding(unwrapped_model, False) + + self._run_dummy_serial_mtp_forward() + + def _run_dummy_async_sched_step(self, input_ids: Tensor, position_ids: Tensor) -> None: + """Run an async-scheduling dummy step in MTP then base-forward order. + + Args: + input_ids (Tensor): Dummy input token IDs. + position_ids (Tensor): Dummy input position IDs. + """ + context = self.inference_wrapped_model.inference_context + if self.model_config.moe_pad_experts_for_cuda_graph_inference: + if not context.using_cuda_graph_this_step(): + set_decode_expert_padding(self._unwrapped_model, False) + + self._run_dummy_serial_mtp_forward() + self._run_dummy_base_forward(input_ids, position_ids) + + @torch.inference_mode() + def dummy_forward(self) -> None: + """Run the mode-specific dummy step used by idle expert-parallel ranks.""" + context = self.inference_wrapped_model.inference_context + input_ids, position_ids, _ = self._dynamic_step_context_init(is_dummy_forward=True) + + if context.config.async_sched_mode == AsyncScheduleMode.LEGACY: + self._run_dummy_legacy_step(input_ids, position_ids) + elif context.config.async_sched_mode == AsyncScheduleMode.ASYNC: + self._run_dummy_async_sched_step(input_ids, position_ids) + else: + raise AssertionError( + f"Unexpected async scheduling mode: {context.config.async_sched_mode}" + ) + + # Clear temporary dummy state while preserving reusable prefix state and counters. + context.reset(preserve_prefix_cache=True, preserve_counters=True) + def _transfer_samples_to_cpu(self, active_request_count: int) -> tuple: """Batch GPU-to-CPU transfer of sampled tokens. @@ -1727,6 +1865,27 @@ def _transfer_samples_to_cpu(self, active_request_count: int) -> tuple: sampled_mtp_tokens_cpu = None return sampled_tokens_cpu, sampled_mtp_tokens_cpu + def _apply_stop_word_finished_ids( + self, active_request_ids: Tensor, active_request_mask: Tensor + ) -> None: + """Mark requests whose generated output matched a stop word as finished. + + Args: + active_request_ids (Tensor): IDs for requests active during the current step. + active_request_mask (Tensor): Mask updated in place for requests that remain active. + """ + if self._get_stop_word_finished_ids_callback is None: + return + + request_ids = active_request_ids.tolist() + stop_word_finished_ids = self._get_stop_word_finished_ids_callback(request_ids) + if not stop_word_finished_ids: + return + + for idx, request_id in enumerate(request_ids): + if request_id in stop_word_finished_ids: + active_request_mask[idx] = 0 + def _dynamic_step_context_bookkeeping(self) -> Dict[str, Tensor]: """Update the dynamic inference context after sampling. @@ -1773,15 +1932,8 @@ def _dynamic_step_context_bookkeeping(self) -> Dict[str, Tensor]: != context.active_request_metadata["termination_id"][:active_request_count] ).byte() & torch.less(active_sequence_lengths, max_sequence_lengths).byte() - # Mark requests as finished if they hit stop words - # (detected in previous step's post_process_requests) - if self._get_stop_word_finished_ids_callback is not None: - request_ids_list = active_request_ids.tolist() - stop_word_finished_ids = self._get_stop_word_finished_ids_callback(request_ids_list) - if stop_word_finished_ids: - for idx, request_id in enumerate(request_ids_list): - if request_id in stop_word_finished_ids: - active_request_mask[idx] = 0 + # Apply stop words detected during the previous engine bookkeeping step. + self._apply_stop_word_finished_ids(active_request_ids, active_request_mask) finished_idxs = ( torch.nonzero(active_request_mask == 0, as_tuple=True)[0] + context.paused_request_count @@ -1822,85 +1974,1174 @@ def _dynamic_step_context_bookkeeping(self) -> Dict[str, Tensor]: **(update_result or {}), } - async def _run_legacy_step(self, skip_bookkeeping: Optional[bool] = False) -> Optional[Dict]: - """Forward step the model and update the inference context. + # ------------------------------------------------------------------------- + # Begin async scheduling methods + # ------------------------------------------------------------------------- + + def _validate_async_sched_support_for_step(self, run_async_overlap: bool) -> None: + """Validate controller/context state for async scheduling. Args: - skip_bookkeeping (Optional[bool]): If true, skip the context bookkeeping step. + run_async_overlap (bool): Whether this step uses overlap ordering. - Returns: - (Optional[Dict]): A dictionary containing: - active_request_ids (Tensor): Current active request IDs. - newly_paused_request_ids (Tensor): Newly paused request IDs. - finished_request_ids (Tensor): Finished request IDs. - sample (Tensor): New sample. - log_probs (Optional[Tensor]): Log probabilities of the new sample, if requested. - cuda_graph_request_count (Optional[int]): Size of cuda graph used for this step. + Raises if the current step does not support async scheduling. """ context = self.inference_wrapped_model.inference_context - self._decode_forward_primer.clear() active_request_count = context.total_request_count - context.paused_request_count - - # No tokens and no active requests? if context.active_token_count == 0 and active_request_count == 0: - return None + return - with torch.inference_mode(): - input_ids, position_ids = self._dynamic_step_context_init() + if run_async_overlap and context.paused_request_count != 0: + raise RuntimeError("Async scheduling overlap does not support paused requests.") - cuda_graph_request_count = ( - context.padded_active_request_count - if context.using_cuda_graph_this_step() + def _compact_async_sched_logits(self, survivor_idxs: Tensor) -> None: + """Compact pending logits and sampling metadata into survivor order. + + Args: + survivor_idxs (Tensor): Active-row indices for requests that remain + active after async scheduling. + """ + if survivor_idxs.numel() == 0: + self._async_sched_logits.clear() + return + + tokens_per_request = self.num_speculative_tokens + 1 + pending_token_row_indices = self._async_sched_logits.token_row_indices + + identity_idxs = torch.arange(survivor_idxs.numel(), device=survivor_idxs.device) + if torch.equal(survivor_idxs, identity_idxs): + survivor_token_row_indices = ( + pending_token_row_indices[: survivor_idxs.numel() * tokens_per_request] + if pending_token_row_indices is not None else None ) + self._async_sched_logits.set_pending( + self._async_sched_logits.cuda_graph_request_count, survivor_token_row_indices + ) + return - # Enable routing recording before forward pass if routing replay is enabled - config = self.inference_wrapped_model.model.config - if config.moe_enable_routing_replay: - RouterReplay.set_global_router_replay_action(RouterReplayAction.RECORD) + token_offsets = torch.arange(tokens_per_request, device=survivor_idxs.device) + survivor_token_idxs = ( + survivor_idxs[:, None] * tokens_per_request + token_offsets[None, :] + ).flatten() + survivor_token_idxs_cuda = survivor_token_idxs.to(self._all_logits_cuda.device) + survivor_token_row_indices = ( + pending_token_row_indices[survivor_token_idxs_cuda] + if pending_token_row_indices is not None + else None + ) - # Forward pass produces only base logits. When speculative decoding is - # active, MTP logits are computed serially after verification. - range_push("forward_pass") - self._dynamic_step_forward_logits(input_ids, position_ids) + compacted_logits = self._all_logits_cuda[:, survivor_token_idxs_cuda, :].contiguous() + if self._enable_cuda_graph: + self._all_logits_cuda[:, : survivor_token_idxs.numel(), :].copy_(compacted_logits) + else: + self._all_logits_cuda = compacted_logits - # Commit Mamba intermediate states before update_requests, which - # may swap request indices. The Python lists tracking EOS block IDs - # and intermediate offsets are not swapped along with tensors, so - # commit must run while indices are still valid. - if context.is_hybrid_model and context.mamba_slot_allocator is not None: - context.mamba_slot_allocator.commit_intermediate_states() + context = self.inference_wrapped_model.inference_context + gpu_view = context.gpu_view + survivor_count = survivor_idxs.numel() + survivor_idxs_cpu = survivor_idxs.to("cpu") + survivor_idxs_cuda = survivor_idxs.to(gpu_view.temperature.device) + for label in ("temperature", "top_k", "top_p"): + compacted_metadata = context.active_request_metadata[label][survivor_idxs_cpu] + context.active_request_metadata[label][:survivor_count].copy_(compacted_metadata) + compacted_temperature = gpu_view.temperature[survivor_idxs_cuda].contiguous() + compacted_top_k = gpu_view.top_k[survivor_idxs_cuda].contiguous() + compacted_top_p = gpu_view.top_p[survivor_idxs_cuda].contiguous() + gpu_view.temperature[:survivor_count].copy_(compacted_temperature) + gpu_view.top_k[:survivor_count].copy_(compacted_top_k) + gpu_view.top_p[:survivor_count].copy_(compacted_top_p) + + self._async_sched_logits.set_pending( + self._async_sched_logits.cuda_graph_request_count, survivor_token_row_indices + ) - # Collect flat routing indices and scatter them into per-block storage. - # Must be done before update_requests while token-to-block mappings are valid. - # Reconstruction happens from blocks at request completion. - routing_indices = self._router_record_bookkeeping() - context.kv_block_allocator.store_routing_per_block(routing_indices) + @staticmethod + def _synchronize_async_sched_event(event: Optional[torch.cuda.Event]) -> None: + """Block the host until an async-scheduling CUDA event completes. - # Save routing indices. - tracer = get_moe_router_tracer() - if tracer is not None and routing_indices is not None: - layer_ids = [ - r.layer_number - for r in RouterReplay.global_router_replay_instances - if r.layer_number is not None - ] or None - tracer.record_indices(torch.from_numpy(routing_indices), layer_ids=layer_ids) - tracer.advance_step() - range_pop() + Args: + event (Optional[torch.cuda.Event]): CUDA event to synchronize, or + `None` when no CUDA work was recorded. + """ + if event is not None: + event.synchronize() - # This is the best place to yield control back to event loop. - # At this point we have enqueued FW pass GPU kernels asynchronously. - # While they are running, we can do other useful CPU work. - # Note: This can be moved further ahead if sampling can be made - # asynchronous. - # Todo [Siddharth]: Can we condition the sleep on a cuda event? - # NOTE [TDE]: This will be moved once CPU and GPU methods are separated. - await asyncio.sleep(0) + def _copy_async_sched_accepted_counts_to_cpu( + self, accepted_counts_gpu: Tensor + ) -> Tuple[Tensor, Optional[torch.cuda.Event]]: + """Start copying MTP acceptance counts into their reusable CPU buffer. - with torch.inference_mode(): - range_push("sampling") - return_log_probs, return_top_n_logprobs = self._dynamic_step_log_probs_bookkeeping() + Args: + accepted_counts_gpu (Tensor): Accepted MTP draft count per active request. + + Returns: + Tuple[Tensor, Optional[torch.cuda.Event]]: Transient CPU view and its + copy-completion event. + """ + if not accepted_counts_gpu.is_cuda: + return accepted_counts_gpu.cpu(), None + + accepted_counts_cpu = self._async_sched_accepted_counts_cpu_buffer[ + : accepted_counts_gpu.numel() + ] + with torch.cuda.stream(self._async_sched_copy_stream): + self._async_sched_copy_stream.wait_event( + self._async_sched_mtp_verification_gpu_ready_event + ) + accepted_counts_cpu.copy_(accepted_counts_gpu, non_blocking=True) + self._async_sched_accepted_counts_cpu_ready_event.record(self._async_sched_copy_stream) + return accepted_counts_cpu, self._async_sched_accepted_counts_cpu_ready_event + + def _copy_async_sched_sample_to_cpu( + self, + sampled_tokens_gpu: Tensor, + sampled_mtp_tokens_gpu: Optional[Tensor] = None, + accepted_tokens_gpu: Optional[Tensor] = None, + ) -> Tuple[Tensor, Optional[Tensor], Optional[Tensor], Optional[torch.cuda.Event]]: + """Start copying async sampling outputs into reusable CPU buffers. + + Args: + sampled_tokens_gpu (Tensor): Sampled base token IDs for active requests. + sampled_mtp_tokens_gpu (Optional[Tensor]): Generated MTP draft token IDs. + accepted_tokens_gpu (Optional[Tensor]): Accepted pending MTP draft token IDs. + + Returns: + Tuple[Tensor, Optional[Tensor], Optional[Tensor], Optional[torch.cuda.Event]]: + Transient CPU views for base, draft, and accepted tokens plus the + copy-completion event. + """ + if not sampled_tokens_gpu.is_cuda: + return ( + sampled_tokens_gpu.cpu(), + sampled_mtp_tokens_gpu.cpu() if sampled_mtp_tokens_gpu is not None else None, + accepted_tokens_gpu.cpu() if accepted_tokens_gpu is not None else None, + None, + ) + + sample_cpu = self._async_sched_sampled_tokens_cpu_buffer[: sampled_tokens_gpu.numel()] + sampled_mtp_tokens_cpu = None + if sampled_mtp_tokens_gpu is not None: + sampled_mtp_tokens_cpu = self._async_sched_sampled_mtp_tokens_cpu_buffer[ + :, : sampled_tokens_gpu.numel() + ] + accepted_tokens_cpu = None + if accepted_tokens_gpu is not None: + accepted_tokens_cpu = self._async_sched_accepted_tokens_cpu_buffer[ + : sampled_tokens_gpu.numel() + ] + + with torch.cuda.stream(self._async_sched_copy_stream): + self._async_sched_copy_stream.wait_event(self._async_sched_sample_gpu_ready_event) + sample_cpu.copy_(sampled_tokens_gpu, non_blocking=True) + if sampled_mtp_tokens_gpu is not None: + sampled_mtp_tokens_cpu.copy_(sampled_mtp_tokens_gpu, non_blocking=True) + if accepted_tokens_gpu is not None: + accepted_tokens_cpu.copy_(accepted_tokens_gpu, non_blocking=True) + self._async_sched_sample_cpu_ready_event.record(self._async_sched_copy_stream) + return ( + sample_cpu, + sampled_mtp_tokens_cpu, + accepted_tokens_cpu, + self._async_sched_sample_cpu_ready_event, + ) + + def _build_async_sched_request_state( + self, sampled_tokens_cpu: Tensor, resolved_sequence_lengths: Tensor + ) -> Tuple[Tensor, Tensor, Tensor]: + """Build request IDs and the active/finished mask for resolution. + + Args: + sampled_tokens_cpu (Tensor): Sampled CPU token IDs for active requests. + resolved_sequence_lengths (Tensor): Sequence lengths after accepting + current output and before preparing unverified successor tokens. + + Returns: + Tuple[Tensor, Tensor, Tensor]: Active request IDs, finished request + IDs, and the active-request mask. + """ + context = self.inference_wrapped_model.inference_context + active_request_count = context.total_request_count - context.paused_request_count + active_request_slice = slice(context.paused_request_count, context.total_request_count) + active_request_ids = context.request_ids[active_request_slice].long() + + max_sequence_lengths = context.get_max_sequence_lengths() + active_request_mask = ( + sampled_tokens_cpu != context.request_metadata["termination_id"][active_request_slice] + ).byte() & torch.less(resolved_sequence_lengths, max_sequence_lengths).byte() + + self._apply_stop_word_finished_ids(active_request_ids, active_request_mask) + + if context.chunked_prefill_request_id != -1: + chunked_prefill_rows = torch.nonzero( + active_request_ids == context.chunked_prefill_request_id, as_tuple=True + )[0] + assert ( + chunked_prefill_rows.numel() == 1 + ), "The active chunked-prefill request must have exactly one row." + active_request_mask[chunked_prefill_rows[0]] = 1 + + finished_idxs = ( + torch.nonzero(active_request_mask == 0, as_tuple=True)[0] + context.paused_request_count + ) + finished_request_ids = context.request_ids[finished_idxs].clone() + assert sampled_tokens_cpu.numel() == active_request_count + + return active_request_ids, finished_request_ids, active_request_mask + + def _run_async_sched_sample(self) -> _AsyncScheduleSampleResult: + """Sample active requests and start transferring their tokens to CPU. + + Returns: + _AsyncScheduleSampleResult: Base-token samples and transfer state. + """ + context = self.inference_wrapped_model.inference_context + active_request_count = context.total_request_count - context.paused_request_count + + range_push("sampling") + self._dynamic_step_sample_logits() + sampled_tokens_gpu = self._sampled_tokens_cuda[:active_request_count] + if sampled_tokens_gpu.is_cuda: + self._async_sched_sample_gpu_ready_event.record( + torch.cuda.current_stream(sampled_tokens_gpu.device) + ) + range_pop() + + sampled_tokens_cpu, _, _, sample_cpu_ready_event = self._copy_async_sched_sample_to_cpu( + sampled_tokens_gpu + ) + return _AsyncScheduleSampleResult( + sampled_tokens_gpu=sampled_tokens_gpu, + sampled_tokens_cpu_view=sampled_tokens_cpu, + sampled_mtp_tokens_gpu=None, + sampled_mtp_tokens_cpu_view=None, + accepted_tokens_cpu_view=None, + accepted_counts_gpu=None, + accepted_counts_cpu_view=None, + accepted_counts_cpu_ready_event=None, + sample_cpu_ready_event=sample_cpu_ready_event, + ) + + def _run_async_sched_sample_mtp(self) -> _AsyncScheduleSampleResult: + """Verify pending MTP logits and generate the next draft tokens. + + Returns: + _AsyncScheduleSampleResult: Base, draft, accepted-token, and transfer state. + """ + context = self.inference_wrapped_model.inference_context + active_request_count = context.total_request_count - context.paused_request_count + token_row_indices = self._async_sched_logits.token_row_indices + if token_row_indices is None: + raise RuntimeError("Pending async MTP logits are missing token-row indices.") + + range_push("sampling") + pending_input_ids = context.gpu_view.token_to_input_ids.unsqueeze(0) + self._dynamic_step_sample_logits_and_verify_tokens( + pending_input_ids, token_row_indices=token_row_indices + ) + accepted_counts_gpu = self._accepted_token_counts_per_request[:active_request_count] + if accepted_counts_gpu.is_cuda: + self._async_sched_mtp_verification_gpu_ready_event.record( + torch.cuda.current_stream(accepted_counts_gpu.device) + ) + accepted_counts_cpu, accepted_counts_cpu_ready_event = ( + self._copy_async_sched_accepted_counts_to_cpu(accepted_counts_gpu) + ) + + base_position = (context.gpu_view.token_to_pos_ids[self._last_accepted_seq_indices] + 1).to( + torch.int64 + ) + self._compute_serial_mtp_and_sample(base_position=base_position) + sampled_tokens_gpu = self._sampled_tokens_cuda[:active_request_count] + sampled_mtp_tokens_gpu = self._sampled_mtp_tokens_cuda[:, :active_request_count] + accepted_tokens_gpu = ( + self._accepted_tokens_per_request[:active_request_count] + if context.num_decode_requests > 0 + else None + ) + if sampled_tokens_gpu.is_cuda: + self._async_sched_sample_gpu_ready_event.record( + torch.cuda.current_stream(sampled_tokens_gpu.device) + ) + range_pop() + + sampled_tokens_cpu, sampled_mtp_tokens_cpu, accepted_tokens_cpu, sample_cpu_ready_event = ( + self._copy_async_sched_sample_to_cpu( + sampled_tokens_gpu, sampled_mtp_tokens_gpu, accepted_tokens_gpu + ) + ) + return _AsyncScheduleSampleResult( + sampled_tokens_gpu=sampled_tokens_gpu, + sampled_tokens_cpu_view=sampled_tokens_cpu, + sampled_mtp_tokens_gpu=sampled_mtp_tokens_gpu, + sampled_mtp_tokens_cpu_view=sampled_mtp_tokens_cpu, + accepted_tokens_cpu_view=accepted_tokens_cpu, + accepted_counts_gpu=accepted_counts_gpu, + accepted_counts_cpu_view=accepted_counts_cpu, + accepted_counts_cpu_ready_event=accepted_counts_cpu_ready_event, + sample_cpu_ready_event=sample_cpu_ready_event, + ) + + def _run_async_sched_mtp_rewind(self, sample_result: _AsyncScheduleSampleResult) -> None: + """Rewind rejected MTP KV state before preparing the successor. + + Args: + sample_result (_AsyncScheduleSampleResult): Verified MTP sampling state. + """ + accepted_counts_cpu = sample_result.accepted_counts_cpu_view + if accepted_counts_cpu is None: + raise RuntimeError("Async MTP sampling did not produce accepted-token counts.") + + self._synchronize_async_sched_event(sample_result.accepted_counts_cpu_ready_event) + + blocks_to_release, remove_mask = self._rewind_kv_cache(accepted_counts_cpu) + context = self.inference_wrapped_model.inference_context + context.kv_block_allocator.release_memory_blocks(blocks_to_release[remove_mask]) + + def _run_async_sched_log_probs( + self, sample_result: _AsyncScheduleSampleResult + ) -> Optional[_AsyncScheduleLogProbsGPUResult]: + """Calculate selected and top-n log probabilities on the GPU. + + Args: + sample_result (_AsyncScheduleSampleResult): Sampled and accepted + tokens for active requests. + + Returns: + Optional[_AsyncScheduleLogProbsGPUResult]: GPU logprob outputs and + their completion event, or `None` when no request needs logprobs. + """ + return_log_probs, return_top_n_logprobs = self._dynamic_step_log_probs_bookkeeping() + if not return_log_probs and not return_top_n_logprobs: + return None + + context = self.inference_wrapped_model.inference_context + active_request_count = context.total_request_count - context.paused_request_count + num_decode_requests = context.num_decode_requests + num_prefill_requests = active_request_count - num_decode_requests + tokens_per_request = self.num_speculative_tokens + 1 + sampled_tokens_gpu = sample_result.sampled_tokens_gpu + + if context.config.materialize_only_last_token_logits: + prefill_row_counts = [1] * num_prefill_requests + else: + active_slice = slice(context.paused_request_count, context.total_request_count) + prefill_row_counts = context.request_query_lengths[ + active_slice.start + num_decode_requests : active_slice.stop + ].tolist() + row_counts = [tokens_per_request] * num_decode_requests + prefill_row_counts + + if self.num_speculative_tokens == 0: + selected_log_probs, log_probs = context.calculate_log_probs_tensors( + self._all_logits_cuda, + sampled_tokens_gpu, + only_last_token_logits=context.config.materialize_only_last_token_logits, + sampling=self._sampling, + ) + else: + accepted_counts_gpu = sample_result.accepted_counts_gpu + if accepted_counts_gpu is None or self._accepted_tokens_per_request is None: + raise RuntimeError("Async MTP sampling did not produce accepted-token state.") + + decode_samples = sampled_tokens_gpu[:num_decode_requests] + decode_tokens = torch.cat( + ( + self._accepted_tokens_per_request[:num_decode_requests].clamp(min=0), + decode_samples.unsqueeze(1), + ), + dim=1, + ) + decode_tokens.scatter_( + 1, + accepted_counts_gpu[:num_decode_requests].unsqueeze(1), + decode_samples.unsqueeze(1), + ) + + if context.config.materialize_only_last_token_logits: + prefill_tokens = sampled_tokens_gpu[num_decode_requests:] + else: + decode_token_count = num_decode_requests * tokens_per_request + prefill_tokens = context.gpu_view.token_to_input_ids[ + decode_token_count : context.active_token_count + ].roll(-1, 0) + prefill_lengths_gpu = context.gpu_view.request_query_lengths[ + num_decode_requests:active_request_count + ] + prefill_last_token_idxs = prefill_lengths_gpu.cumsum(0) - 1 + prefill_tokens[prefill_last_token_idxs] = sampled_tokens_gpu[num_decode_requests:] + + selected_tokens = torch.cat((decode_tokens.flatten(), prefill_tokens)) + logit_count = sum(row_counts) + logits = self._all_logits_cuda[:, :logit_count, :] + row_to_request = torch.arange(active_request_count).repeat_interleave( + torch.tensor(row_counts) + ) + selected_log_probs, log_probs = context.calculate_log_probs_tensors( + logits, + selected_tokens, + only_last_token_logits=True, + sampling=self._sampling, + row_to_request=row_to_request, + ) + + top_n_counts = ( + context.active_request_metadata["top_n_logprobs"][:active_request_count].tolist() + if return_top_n_logprobs + else [0] * active_request_count + ) + skip_prompt_log_probs = context.active_request_metadata["skip_prompt_log_probs"][ + :active_request_count + ].tolist() + max_top_n = max(top_n_counts, default=0) + if max_top_n > 0: + top_n_result = torch.topk(log_probs[: sum(row_counts)], k=max_top_n, dim=-1) + top_n_log_probs = top_n_result.values + top_n_token_ids = top_n_result.indices + else: + top_n_log_probs = None + top_n_token_ids = None + + gpu_ready_event = None + if selected_log_probs.is_cuda: + current_stream = torch.cuda.current_stream(selected_log_probs.device) + self._async_sched_log_probs_gpu_ready_event.record(current_stream) + gpu_ready_event = self._async_sched_log_probs_gpu_ready_event + + return _AsyncScheduleLogProbsGPUResult( + selected_log_probs=selected_log_probs, + top_n_log_probs=top_n_log_probs, + top_n_token_ids=top_n_token_ids, + row_counts=row_counts, + top_n_counts=top_n_counts, + skip_prompt_log_probs=skip_prompt_log_probs, + num_decode_requests=num_decode_requests, + gpu_ready_event=gpu_ready_event, + ) + + def _copy_async_sched_log_probs_to_cpu( + self, gpu_result: Optional[_AsyncScheduleLogProbsGPUResult] + ) -> Optional[_AsyncScheduleLogProbsTransfer]: + """Start selected and top-n logprob transfers to reusable CPU buffers. + + Args: + gpu_result (Optional[_AsyncScheduleLogProbsGPUResult]): GPU outputs + produced by the current sampling step. + + Returns: + Optional[_AsyncScheduleLogProbsTransfer]: Transient CPU views, + transfer-completion event, and retained GPU sources, or `None`. + """ + if gpu_result is None: + return None + + selected_log_probs = gpu_result.selected_log_probs + selected_shape = selected_log_probs.shape + selected_size = selected_log_probs.numel() + max_top_n = max(gpu_result.top_n_counts, default=0) + if max_top_n > self._async_sched_top_n_capacity: + context = self.inference_wrapped_model.inference_context + buffer_size = context.max_tokens * max_top_n + self._async_sched_top_n_log_probs_cpu_buffer = torch.empty( + buffer_size, dtype=torch.float32, device="cpu", pin_memory=True + ) + self._async_sched_top_n_token_ids_cpu_buffer = torch.empty( + buffer_size, dtype=torch.int64, device="cpu", pin_memory=True + ) + self._async_sched_top_n_capacity = max_top_n + + selected_log_probs_cpu_view = self._async_sched_selected_log_probs_cpu_buffer[ + :selected_size + ].view(selected_shape) + if max_top_n > 0: + top_n_size = selected_size * max_top_n + top_n_log_probs_cpu_view = self._async_sched_top_n_log_probs_cpu_buffer[ + :top_n_size + ].view(*selected_shape, max_top_n) + top_n_token_ids_cpu_view = self._async_sched_top_n_token_ids_cpu_buffer[ + :top_n_size + ].view(*selected_shape, max_top_n) + else: + top_n_log_probs_cpu_view = None + top_n_token_ids_cpu_view = None + + cpu_ready_event = None + if selected_log_probs.is_cuda: + assert gpu_result.gpu_ready_event is not None + with torch.cuda.stream(self._async_sched_copy_stream): + self._async_sched_copy_stream.wait_event(gpu_result.gpu_ready_event) + selected_log_probs_cpu_view.copy_(selected_log_probs, non_blocking=True) + if max_top_n > 0: + assert gpu_result.top_n_log_probs is not None + assert gpu_result.top_n_token_ids is not None + assert top_n_log_probs_cpu_view is not None + assert top_n_token_ids_cpu_view is not None + top_n_log_probs_cpu_view.copy_(gpu_result.top_n_log_probs, non_blocking=True) + top_n_token_ids_cpu_view.copy_(gpu_result.top_n_token_ids, non_blocking=True) + self._async_sched_log_probs_cpu_ready_event.record(self._async_sched_copy_stream) + cpu_ready_event = self._async_sched_log_probs_cpu_ready_event + else: + selected_log_probs_cpu_view.copy_(selected_log_probs) + if max_top_n > 0: + assert gpu_result.top_n_log_probs is not None + assert gpu_result.top_n_token_ids is not None + assert top_n_log_probs_cpu_view is not None + assert top_n_token_ids_cpu_view is not None + top_n_log_probs_cpu_view.copy_(gpu_result.top_n_log_probs) + top_n_token_ids_cpu_view.copy_(gpu_result.top_n_token_ids) + + return _AsyncScheduleLogProbsTransfer( + selected_log_probs_cpu_view=selected_log_probs_cpu_view, + top_n_log_probs_cpu_view=top_n_log_probs_cpu_view, + top_n_token_ids_cpu_view=top_n_token_ids_cpu_view, + row_counts=gpu_result.row_counts, + top_n_counts=gpu_result.top_n_counts, + skip_prompt_log_probs=gpu_result.skip_prompt_log_probs, + num_decode_requests=gpu_result.num_decode_requests, + cpu_ready_event=cpu_ready_event, + gpu_result=gpu_result, + ) + + @staticmethod + def _materialize_async_sched_log_probs( + transfer: Optional[_AsyncScheduleLogProbsTransfer], + accepted_counts_cpu: Optional[Tensor] = None, + ) -> Tuple[Optional[List[List[float]]], Optional[Dict[int, List[Tuple[Tensor, Tensor]]]]]: + """Convert completed CPU transfer views to the legacy result format. + + Args: + transfer (Optional[_AsyncScheduleLogProbsTransfer]): Completed + logprob transfer for the current step. + accepted_counts_cpu (Optional[Tensor]): Accepted MTP draft count per + active request, or `None` for one-token decoding. + + Returns: + Tuple containing selected logprobs per request and optional top-n + values/token IDs per request. + """ + if transfer is None: + return None, None + + accepted_counts = accepted_counts_cpu.tolist() if accepted_counts_cpu is not None else None + row_offset = 0 + log_probs = [] + top_n_logprobs = {} + for request_idx, (row_count, top_n, skip_prompt) in enumerate( + zip(transfer.row_counts, transfer.top_n_counts, transfer.skip_prompt_log_probs) + ): + is_decode = request_idx < transfer.num_decode_requests + emitted_count = ( + accepted_counts[request_idx] + 1 + if is_decode and accepted_counts is not None + else row_count + ) + emitted_slice = slice(row_offset, row_offset + emitted_count) + log_probs.append(transfer.selected_log_probs_cpu_view[emitted_slice].tolist()) + + if top_n > 0: + assert transfer.top_n_log_probs_cpu_view is not None + assert transfer.top_n_token_ids_cpu_view is not None + if not is_decode and skip_prompt: + top_n_row_idxs = [row_offset + row_count - 1] + else: + top_n_row_idxs = range(row_offset, row_offset + emitted_count) + top_n_logprobs[request_idx] = [ + ( + transfer.top_n_log_probs_cpu_view[token_idx, :top_n].clone(), + transfer.top_n_token_ids_cpu_view[token_idx, :top_n].clone(), + ) + for token_idx in top_n_row_idxs + ] + row_offset += row_count + + return log_probs, top_n_logprobs or None + + def _run_async_sched_prepare(self) -> Tuple[Tensor, Tensor]: + """Prepare decode requests and return live GPU forward-input views. + + The returned views have their final shape and stable backing storage, + but their contents are populated later. Sampling updates the input-ID + view, and deferred bookkeeping publication updates the position-ID view. + + Returns: + Tuple[Tensor, Tensor]: Live GPU input-ID and position-ID views for + the speculative forward. + """ + context = self.inference_wrapped_model.inference_context + context.prepare_requests() + input_ids, position_ids, _ = self._dynamic_step_context_init( + transfer_bookkeeping_to_gpu=False + ) + return input_ids, position_ids + + def _run_async_sched_publish_bookkeeping(self) -> Optional[torch.cuda.Event]: + """Publish prepared bookkeeping without overwriting GPU input token IDs. + + Returns: + Optional[torch.cuda.Event]: Event marking bookkeeping H2D completion. + """ + context = self.inference_wrapped_model.inference_context + return context.transfer_bookkeeping_to_gpu( + skip_token_input_ids=True, record_done_event=True + ) + + def _commit_mamba_intermediate_states(self) -> None: + """Commit prefix-cacheable Mamba states produced by the current forward.""" + context = self.inference_wrapped_model.inference_context + if context.is_hybrid_model and context.mamba_slot_allocator is not None: + context.mamba_slot_allocator.commit_intermediate_states() + + def _run_async_sched_forward( + self, input_ids_gpu_view: Tensor, position_ids_gpu_view: Tensor + ) -> None: + """Run one dynamic forward pass and cache logits for async scheduling. + + Args: + input_ids_gpu_view (Tensor): Live GPU view of the input token IDs. + position_ids_gpu_view (Tensor): Live GPU view of the position IDs. + """ + context = self.inference_wrapped_model.inference_context + cuda_graph_request_count = ( + context.padded_active_request_count if context.using_cuda_graph_this_step() else None + ) + + # Forward. + range_push("forward_pass") + self._dynamic_step_forward_logits(input_ids_gpu_view, position_ids_gpu_view) + self._commit_mamba_intermediate_states() + range_pop() + + # Record the logits and identity mapping for this forward's input rows. + token_row_indices = None + if self._async_sched_mtp_token_row_indices is not None: + token_row_indices = self._async_sched_mtp_token_row_indices[ + : context.active_token_count + ] + self._async_sched_logits.set_pending(cuda_graph_request_count, token_row_indices) + + def _run_dummy_async_sched_base_step(self) -> None: + """Run the base-forward half of an async EP step after local work finishes.""" + context = self.inference_wrapped_model.inference_context + + input_ids, position_ids, _ = self._dynamic_step_context_init(is_dummy_forward=True) + self._run_dummy_base_forward(input_ids, position_ids) + context.reset(preserve_prefix_cache=True, preserve_counters=True) + + def _run_async_sched_forward_primer(self) -> Tuple[bool, Optional[torch.cuda.Event]]: + """Launch the initial forward when no valid logits state exists. + + Returns: + Tuple[bool, Optional[torch.cuda.Event]]: Whether this call launched + the forward primer and its bookkeeping H2D completion event. + """ + if self._async_sched_logits.is_valid: + return False, None + + # Initialize, forward, and record the pending logits state. + with torch.inference_mode(): + input_ids_gpu_view, position_ids_gpu_view, bookkeeping_done_event = ( + self._dynamic_step_context_init(record_bookkeeping_done_event=True) + ) + if self.num_speculative_tokens > 0 and self.model_config.expert_model_parallel_size > 1: + self._run_dummy_serial_mtp_forward() + self._run_async_sched_forward(input_ids_gpu_view, position_ids_gpu_view) + + return True, bookkeeping_done_event + + def _run_async_sched_resolve( + self, sample_result: _AsyncScheduleSampleResult, resolved_sequence_lengths: Tensor + ) -> _AsyncScheduleRequestResult: + """Resolve request state and compact speculative forward logits. + + Args: + sample_result (_AsyncScheduleSampleResult): Sampling outputs in reusable CPU views. + resolved_sequence_lengths (Tensor): Sequence lengths after accepting + current output and before preparing unverified successor tokens. + + Returns: + _AsyncScheduleRequestResult: Sampled tokens, resolved request row + sets, and survivor indices. + """ + context = self.inference_wrapped_model.inference_context + + # Clone the transient D2H view before the next step can reuse its buffer. + range_push("active_request_mask") + sampled_tokens_cpu = sample_result.sampled_tokens_cpu_view.clone() + accepted_tokens_cpu = ( + sample_result.accepted_tokens_cpu_view.clone() + if sample_result.accepted_tokens_cpu_view is not None + else None + ) + active_request_ids, finished_request_ids, active_request_mask = ( + self._build_async_sched_request_state(sampled_tokens_cpu, resolved_sequence_lengths) + ) + range_pop() + + # Resolve CPU request lifecycle state. + range_push("resolve_requests") + resolved_finished_request_ids, survivor_idxs = context.resolve_requests(active_request_mask) + range_pop() + + assert torch.equal(finished_request_ids, resolved_finished_request_ids) + + # Enqueue compaction behind the successor forward on the current CUDA stream. + self._compact_async_sched_logits(survivor_idxs) + + # Return the resolution result. + return _AsyncScheduleRequestResult( + sampled_tokens_cpu=sampled_tokens_cpu, + accepted_tokens_cpu=accepted_tokens_cpu, + active_request_ids=active_request_ids, + finished_request_ids=finished_request_ids, + survivor_idxs=survivor_idxs, + ) + + def _run_async_sched_update_requests( + self, sample_result: _AsyncScheduleSampleResult, resolved_sequence_lengths: Tensor + ) -> _AsyncScheduleRequestResult: + """Run complete request lifecycle bookkeeping for a no-overlap step. + + Args: + sample_result (_AsyncScheduleSampleResult): Sampling outputs in reusable CPU views. + resolved_sequence_lengths (Tensor): Sequence lengths after accepting + the current output. + + Returns: + _AsyncScheduleRequestResult: Stable sampled output and lifecycle results. + """ + context = self.inference_wrapped_model.inference_context + + sampled_tokens_cpu = sample_result.sampled_tokens_cpu_view.clone() + accepted_tokens_cpu = ( + sample_result.accepted_tokens_cpu_view.clone() + if sample_result.accepted_tokens_cpu_view is not None + else None + ) + active_request_ids, finished_request_ids, active_request_mask = ( + self._build_async_sched_request_state(sampled_tokens_cpu, resolved_sequence_lengths) + ) + + mutable_sampled_tokens_cpu = sampled_tokens_cpu.clone() + mutable_sampled_mtp_tokens_cpu = ( + sample_result.sampled_mtp_tokens_cpu_view.clone() + if sample_result.sampled_mtp_tokens_cpu_view is not None + else None + ) + + range_push("update_requests") + update_result = context.update_requests( + active_request_mask, mutable_sampled_tokens_cpu, mutable_sampled_mtp_tokens_cpu + ) + range_pop() + update_result = update_result or {} + + return _AsyncScheduleRequestResult( + sampled_tokens_cpu=sampled_tokens_cpu, + accepted_tokens_cpu=accepted_tokens_cpu, + active_request_ids=active_request_ids, + finished_request_ids=finished_request_ids, + newly_paused_request_ids=update_result.get("newly_paused_request_ids"), + evict_request_ids=update_result.get("evict_request_ids"), + ) + + def _build_async_sched_step_result( + self, + request_result: _AsyncScheduleRequestResult, + cuda_graph_request_count: Optional[int], + decode_only: DecodeOnly, + log_probs: Optional[List[List[float]]], + top_n_logprobs: Optional[Dict[int, List[Tuple[Tensor, Tensor]]]], + *, + count_compaction: bool, + ) -> DynamicBatchControllerStepResult: + """Build the public result and update async-scheduling counters. + + Args: + request_result (_AsyncScheduleRequestResult): Completed request bookkeeping. + cuda_graph_request_count (Optional[int]): CUDA graph request count used + by the consumed forward. + decode_only (DecodeOnly): Decode-only state for the consumed and + launched forwards. + log_probs (Optional[List[List[float]]]): Selected-token log probabilities + grouped by active request. + top_n_logprobs (Optional[Dict[int, List[Tuple[Tensor, Tensor]]]]): Top-n + log probabilities and token IDs grouped by active request. + count_compaction (bool): Whether finished requests discarded successor rows. + + Returns: + DynamicBatchControllerStepResult: Completed sampled-step result. + """ + context = self.inference_wrapped_model.inference_context + context.async_sched_step_count += 1 + if count_compaction and request_result.finished_request_ids.numel() > 0: + context.async_sched_compaction_step_count += 1 + + return DynamicBatchControllerStepResult( + decode_only=decode_only, + output={ + "active_request_ids": request_result.active_request_ids, + "finished_request_ids": request_result.finished_request_ids, + "sample": request_result.sampled_tokens_cpu, + "finished_routing_block_ids": {}, + "newly_paused_request_ids": request_result.newly_paused_request_ids, + "evict_request_ids": request_result.evict_request_ids, + "accepted_tokens": request_result.accepted_tokens_cpu, + "log_probs": log_probs, + "top_n_logprobs": top_n_logprobs, + "cuda_graph_request_count": cuda_graph_request_count, + }, + ) + + async def _run_async_sched_step_no_overlap( + self, *, schedule_waiting_requests: Optional[Callable[[], None]] + ) -> DynamicBatchControllerStepResult: + """Run ``sample/MTP -> update -> admit -> forward``. + + The first call in an active chain has no pending output. It skips the + first two phases, admits requests, and launches a primer-only forward. + + Args: + schedule_waiting_requests (Optional[Callable[[], None]]): Engine callback + that admits eligible non-chunked prefill requests. + + Returns: + DynamicBatchControllerStepResult: Primer-only state or sampled output. + """ + context = self.inference_wrapped_model.inference_context + had_pending_forward = self._async_sched_logits.is_valid + consumed_decode_only = context.is_decode_only() if had_pending_forward else None + launched_decode_only = None + request_result = None + cuda_graph_request_count = None + log_probs_transfer = None + + with torch.inference_mode(): + if had_pending_forward: + cuda_graph_request_count = self._async_sched_logits.cuda_graph_request_count + + # ------------------------------------------------------------------------- + # Sample/MTP + # ------------------------------------------------------------------------- + if self.num_speculative_tokens > 0: + sample_result = self._run_async_sched_sample_mtp() + self._run_async_sched_mtp_rewind(sample_result) + else: + sample_result = self._run_async_sched_sample() + + log_probs_gpu_result = self._run_async_sched_log_probs(sample_result) + log_probs_transfer = self._copy_async_sched_log_probs_to_cpu(log_probs_gpu_result) + + self._synchronize_async_sched_event(sample_result.sample_cpu_ready_event) + + # ------------------------------------------------------------------------- + # Update + # ------------------------------------------------------------------------- + resolved_sequence_lengths = context.get_active_sequence_lengths() + 1 + + self._async_sched_logits.clear() + request_result = self._run_async_sched_update_requests( + sample_result, resolved_sequence_lengths + ) + + # ------------------------------------------------------------------------- + # Admit + # ------------------------------------------------------------------------- + # This is the only async-scheduling admission mutation point. + if schedule_waiting_requests is not None: + schedule_waiting_requests() + + # ------------------------------------------------------------------------- + # Forward + # ------------------------------------------------------------------------- + active_request_count = context.total_request_count - context.paused_request_count + if active_request_count > 0: + if had_pending_forward: + input_ids, position_ids, _ = self._dynamic_step_context_init() + launched_decode_only = context.is_decode_only() + self._run_async_sched_forward(input_ids, position_ids) + else: + primer_launched, bookkeeping_done_event = self._run_async_sched_forward_primer() + assert primer_launched, "Initial no-overlap step must launch a forward primer." + launched_decode_only = context.is_decode_only() + self._synchronize_async_sched_event(bookkeeping_done_event) + elif had_pending_forward and self.model_config.expert_model_parallel_size > 1: + self._run_dummy_async_sched_base_step() + + decode_only = DecodeOnly(consumed=consumed_decode_only, launched=launched_decode_only) + if not had_pending_forward: + assert active_request_count > 0, "Async no-overlap admission did not add a request." + return DynamicBatchControllerStepResult(decode_only=decode_only, primer_only=True) + + if log_probs_transfer is not None: + self._synchronize_async_sched_event(log_probs_transfer.cpu_ready_event) + log_probs, top_n_logprobs = self._materialize_async_sched_log_probs( + log_probs_transfer, + sample_result.accepted_counts_cpu_view if self.num_speculative_tokens > 0 else None, + ) + result = self._build_async_sched_step_result( + request_result, + cuda_graph_request_count, + decode_only, + log_probs, + top_n_logprobs, + count_compaction=False, + ) + await asyncio.sleep(0) + return result + + async def _run_async_sched_step_overlap(self) -> DynamicBatchControllerStepResult: + """Run ``prepare -> sample -> forward -> resolve`` with one token per request. + + Returns: + DynamicBatchControllerStepResult: Completed sampled-step result. + """ + context = self.inference_wrapped_model.inference_context + assert self._async_sched_logits.is_valid, "Async overlap requires pending logits." + consumed_decode_only = context.is_decode_only() + + with torch.inference_mode(): + cuda_graph_request_count = self._async_sched_logits.cuda_graph_request_count + + resolved_sequence_lengths = context.get_active_sequence_lengths() + 1 + + # ------------------------------------------------------------------------- + # Prepare + # ------------------------------------------------------------------------- + # Prepare CPU state and live GPU views without publishing bookkeeping yet. + range_push("prepare_requests") + input_ids_gpu_view, position_ids_gpu_view = self._run_async_sched_prepare() + range_pop() + launched_decode_only = context.is_decode_only() + decode_only = DecodeOnly(consumed=consumed_decode_only, launched=launched_decode_only) + assert ( + consumed_decode_only and launched_decode_only + ), "Async overlap requires decode-only consumed and launched work." + + # ------------------------------------------------------------------------- + # Sample + # ------------------------------------------------------------------------- + # Enqueue sampling behind the current logits-producing work. + sample_result = self._run_async_sched_sample() + + # Populate the next forward's input-ID view directly from GPU samples. + context.copy_async_sched_sample_to_forward(sample_result.sampled_tokens_gpu) + + log_probs_gpu_result = self._run_async_sched_log_probs(sample_result) + log_probs_transfer = self._copy_async_sched_log_probs_to_cpu(log_probs_gpu_result) + + # ------------------------------------------------------------------------- + # Forward + # ------------------------------------------------------------------------- + # Publish positions and metadata without overwriting GPU-resident input IDs. + range_push("async_sched_transfer_bookkeeping_to_gpu") + bookkeeping_done_event = self._run_async_sched_publish_bookkeeping() + range_pop() + + range_push("async_sched_forward_pass") + self._run_async_sched_forward(input_ids_gpu_view, position_ids_gpu_view) + range_pop() + + # ------------------------------------------------------------------------- + # Resolve + # ------------------------------------------------------------------------- + # Wait for the CPU sample and the published bookkeeping snapshot. + self._synchronize_async_sched_event(sample_result.sample_cpu_ready_event) + self._synchronize_async_sched_event(bookkeeping_done_event) + + # Resolve N while forward N+1 continues. + resolve_result = self._run_async_sched_resolve(sample_result, resolved_sequence_lengths) + + # Commit CPU input IDs in the resolved survivor order. + context.commit_sampled_tokens( + resolve_result.sampled_tokens_cpu[resolve_result.survivor_idxs] + ) + + if log_probs_transfer is not None: + self._synchronize_async_sched_event(log_probs_transfer.cpu_ready_event) + log_probs, top_n_logprobs = self._materialize_async_sched_log_probs(log_probs_transfer) + result = self._build_async_sched_step_result( + resolve_result, + cuda_graph_request_count, + decode_only, + log_probs, + top_n_logprobs, + count_compaction=True, + ) + + # Yield only after resolution is complete and forward N+1 is already submitted. + await asyncio.sleep(0) + return result + + async def _run_async_sched_step_overlap_mtp(self) -> DynamicBatchControllerStepResult: + """Run ``sample/MTP -> prepare -> forward -> resolve`` with MTP. + + Returns: + DynamicBatchControllerStepResult: Completed sampled-step result. + """ + context = self.inference_wrapped_model.inference_context + assert self._async_sched_logits.is_valid, "Async MTP overlap requires pending logits." + consumed_decode_only = context.is_decode_only() + + with torch.inference_mode(): + cuda_graph_request_count = self._async_sched_logits.cuda_graph_request_count + + # ------------------------------------------------------------------------- + # Sample/MTP + # ------------------------------------------------------------------------- + # Verify pending drafts, sample replacements, and rewind rejected KV state. + sample_result = self._run_async_sched_sample_mtp() + self._run_async_sched_mtp_rewind(sample_result) + resolved_sequence_lengths = context.get_active_sequence_lengths() + 1 + log_probs_gpu_result = self._run_async_sched_log_probs(sample_result) + log_probs_transfer = self._copy_async_sched_log_probs_to_cpu(log_probs_gpu_result) + + # ------------------------------------------------------------------------- + # Prepare + # ------------------------------------------------------------------------- + # Prepare CPU state and live GPU views using the verified sequence lengths. + range_push("prepare_requests") + input_ids_gpu_view, position_ids_gpu_view = self._run_async_sched_prepare() + range_pop() + launched_decode_only = context.is_decode_only() + decode_only = DecodeOnly(consumed=consumed_decode_only, launched=launched_decode_only) + assert ( + consumed_decode_only and launched_decode_only + ), "Async MTP overlap requires decode-only consumed and launched work." + + # Populate the next forward with the sampled base and draft tokens. + context.copy_async_sched_sample_to_forward( + sample_result.sampled_tokens_gpu, sample_result.sampled_mtp_tokens_gpu + ) + + # ------------------------------------------------------------------------- + # Forward + # ------------------------------------------------------------------------- + # Publish positions and metadata without overwriting GPU-resident input IDs. + range_push("async_sched_transfer_bookkeeping_to_gpu") + bookkeeping_done_event = self._run_async_sched_publish_bookkeeping() + range_pop() + + range_push("async_sched_forward_pass") + self._run_async_sched_forward(input_ids_gpu_view, position_ids_gpu_view) + range_pop() + + # ------------------------------------------------------------------------- + # Resolve + # ------------------------------------------------------------------------- + # Wait for the CPU samples and the published bookkeeping snapshot. + self._synchronize_async_sched_event(sample_result.sample_cpu_ready_event) + self._synchronize_async_sched_event(bookkeeping_done_event) + + resolve_result = self._run_async_sched_resolve(sample_result, resolved_sequence_lengths) + + # Commit CPU input IDs in the resolved survivor order. + survivor_idxs = resolve_result.survivor_idxs + sampled_mtp_tokens_cpu = ( + sample_result.sampled_mtp_tokens_cpu_view[:, survivor_idxs] + if sample_result.sampled_mtp_tokens_cpu_view is not None + else None + ) + context.commit_sampled_tokens( + resolve_result.sampled_tokens_cpu[survivor_idxs], sampled_mtp_tokens_cpu + ) + + if log_probs_transfer is not None: + self._synchronize_async_sched_event(log_probs_transfer.cpu_ready_event) + log_probs, top_n_logprobs = self._materialize_async_sched_log_probs( + log_probs_transfer, sample_result.accepted_counts_cpu_view + ) + result = self._build_async_sched_step_result( + resolve_result, + cuda_graph_request_count, + decode_only, + log_probs, + top_n_logprobs, + count_compaction=True, + ) + await asyncio.sleep(0) + return result + + # ------------------------------------------------------------------------- + # End async scheduling methods + # ------------------------------------------------------------------------- + + async def _run_legacy_step( + self, skip_bookkeeping: Optional[bool] = False + ) -> DynamicBatchControllerStepResult: + """Forward step the model and update the inference context. + + Args: + skip_bookkeeping (Optional[bool]): If true, skip the context bookkeeping step. + + Returns: + DynamicBatchControllerStepResult: Legacy sampled-step output and its + decode-only state. + """ + context = self.inference_wrapped_model.inference_context + self._async_sched_logits.clear() + active_request_count = context.total_request_count - context.paused_request_count + + # No tokens and no active requests? + if context.active_token_count == 0 and active_request_count == 0: + return DynamicBatchControllerStepResult( + decode_only=DecodeOnly(consumed=None, launched=None) + ) + + with torch.inference_mode(): + input_ids, position_ids, _ = self._dynamic_step_context_init() + is_decode_only = context.is_decode_only() + + cuda_graph_request_count = ( + context.padded_active_request_count + if context.using_cuda_graph_this_step() + else None + ) + + # Enable routing recording before forward pass if routing replay is enabled + config = self.inference_wrapped_model.model.config + if config.moe_enable_routing_replay: + RouterReplay.set_global_router_replay_action(RouterReplayAction.RECORD) + + # Forward pass produces only base logits. When speculative decoding is + # active, MTP logits are computed serially after verification. + range_push("forward_pass") + self._dynamic_step_forward_logits(input_ids, position_ids) + + # Commit Mamba intermediate states before update_requests, which + # may swap request indices. The Python lists tracking EOS block IDs + # and intermediate offsets are not swapped along with tensors, so + # commit must run while indices are still valid. + self._commit_mamba_intermediate_states() + + # Collect flat routing indices and scatter them into per-block storage. + # Must be done before update_requests while token-to-block mappings are valid. + # Reconstruction happens from blocks at request completion. + routing_indices = self._router_record_bookkeeping() + context.kv_block_allocator.store_routing_per_block(routing_indices) + + # Save routing indices. + tracer = get_moe_router_tracer() + if tracer is not None and routing_indices is not None: + layer_ids = [ + r.layer_number + for r in RouterReplay.global_router_replay_instances + if r.layer_number is not None + ] or None + tracer.record_indices(torch.from_numpy(routing_indices), layer_ids=layer_ids) + tracer.advance_step() + range_pop() + + # This is the best place to yield control back to event loop. + # At this point we have enqueued FW pass GPU kernels asynchronously. + # While they are running, we can do other useful CPU work. + # Note: This can be moved further ahead if sampling can be made + # asynchronous. + # Todo [Siddharth]: Can we condition the sleep on a cuda event? + # NOTE [TDE]: This will be moved once CPU and GPU methods are separated. + await asyncio.sleep(0) + + with torch.inference_mode(): + range_push("sampling") + return_log_probs, return_top_n_logprobs = self._dynamic_step_log_probs_bookkeeping() if self.num_speculative_tokens > 0: # Phase 1: Verify speculative tokens using base logits only. @@ -1983,125 +3224,78 @@ async def _run_legacy_step(self, skip_bookkeeping: Optional[bool] = False) -> Op self._accepted_tokens_per_request.fill_(-1) self._accepted_token_counts_per_request.fill_(0) ret.update(request_bookkeeping) - return ret - - async def _run_async_sched_serial_step(self) -> Optional[Dict]: - """Run one decode-only step using serial async scheduling. - - Returns: - Optional[Dict]: Step result for sampled and finished requests, or - `None` when no requests are active. - """ - context = self.inference_wrapped_model.inference_context - active_request_count = context.total_request_count - context.paused_request_count - - if context.active_token_count == 0 and active_request_count == 0: - self._decode_forward_primer.clear() - return None - - self._validate_async_sched_support_for_step() - - with torch.inference_mode(): - if not self._decode_forward_primer.is_primed: - input_ids, position_ids = self._dynamic_step_context_init() - self._run_async_sched_forward(input_ids, position_ids) - - await asyncio.sleep(0) - - with torch.inference_mode(): - active_request_count = context.total_request_count - context.paused_request_count - active_request_slice = slice(context.paused_request_count, context.total_request_count) - active_request_ids = context.request_ids[active_request_slice].long() - - cached_cuda_graph_request_count = self._decode_forward_primer.cuda_graph_request_count - - range_push("sampling") - sampled_tokens_cuda = torch.argmax( - self._all_logits_cuda.squeeze(0)[:active_request_count].float(), dim=-1 + return DynamicBatchControllerStepResult( + decode_only=DecodeOnly(consumed=is_decode_only, launched=is_decode_only), output=ret ) - sampled_tokens_cpu = sampled_tokens_cuda.cpu() - range_pop() - - range_push("active_request_mask") - active_sequence_lengths = context.get_active_sequence_lengths() - active_sequence_lengths += 1 - max_sequence_lengths = context.get_max_sequence_lengths() - active_request_mask = ( - sampled_tokens_cpu - != context.request_metadata["termination_id"][active_request_slice] - ).byte() & torch.less(active_sequence_lengths, max_sequence_lengths).byte() - - finished_idxs = ( - torch.nonzero(active_request_mask == 0, as_tuple=True)[0] - + context.paused_request_count - ) - finished_request_ids = context.request_ids[finished_idxs].clone() - survivor_idxs = torch.nonzero(active_request_mask == 1, as_tuple=True)[0] - new_sample_copy = sampled_tokens_cpu.clone() - range_pop() - - range_push("prepare_requests") - input_ids, position_ids = self._run_async_sched_prepare(new_sample_copy) - range_pop() - - range_push("async_sched_forward_pass") - self._run_async_sched_forward(input_ids, position_ids) - range_pop() - - range_push("resolve_requests") - resolved_finished_request_ids = context.resolve_requests(active_request_mask) - range_pop() - - assert torch.equal(finished_request_ids, resolved_finished_request_ids) - self._compact_async_sched_logits(survivor_idxs) - - context.async_sched_step_count += 1 - if survivor_idxs.numel() < active_request_count: - context.async_sched_compaction_step_count += 1 - - return { - "active_request_ids": active_request_ids, - "finished_request_ids": finished_request_ids, - "sample": sampled_tokens_cpu, - "finished_routing_block_ids": {}, - "newly_paused_request_ids": None, - "evict_request_ids": None, - "accepted_tokens": None, - "log_probs": None, - "top_n_logprobs": None, - "cuda_graph_request_count": cached_cuda_graph_request_count, - } async def async_generate_output_tokens_dynamic_batch( - self, skip_bookkeeping: Optional[bool] = False - ) -> Optional[Dict]: + self, + skip_bookkeeping: Optional[bool] = False, + *, + run_async_overlap: bool = True, + schedule_waiting_requests: Optional[Callable[[], None]] = None, + ) -> DynamicBatchControllerStepResult: """Forward step the model and update the inference context. Args: skip_bookkeeping (Optional[bool]): If true, skip context bookkeeping on the legacy path. + run_async_overlap (bool): Whether to run the overlap ordering. + schedule_waiting_requests (Optional[Callable[[], None]]): Engine callback + used by the no-overlap path to admit eligible prefill requests. Returns: - Optional[Dict]: Step result for sampled and finished requests, or - `None` when no requests are active. + DynamicBatchControllerStepResult: One controller-step result. """ context = self.inference_wrapped_model.inference_context mode = context.config.async_sched_mode - if mode == AsyncScheduleMode.LEGACY or context.num_prefill_requests != 0: + if mode == AsyncScheduleMode.LEGACY: return await self._run_legacy_step(skip_bookkeeping) - if mode == AsyncScheduleMode.SERIAL: - assert not skip_bookkeeping, "Serial async scheduling requires request bookkeeping." - return await self._run_async_sched_serial_step() - raise AssertionError(f"Unexpected async scheduling mode: {mode}") + if mode != AsyncScheduleMode.ASYNC: + raise AssertionError(f"Unexpected async scheduling mode: {mode}") + + assert not skip_bookkeeping, "Async scheduling requires request bookkeeping." + self._validate_async_sched_support_for_step(run_async_overlap) + + active_request_count = context.total_request_count - context.paused_request_count + if context.active_token_count == 0 and active_request_count == 0 and run_async_overlap: + self._async_sched_logits.clear() + return DynamicBatchControllerStepResult( + decode_only=DecodeOnly(consumed=None, launched=None) + ) + + if not run_async_overlap or not self._async_sched_logits.is_valid: + return await self._run_async_sched_step_no_overlap( + schedule_waiting_requests=schedule_waiting_requests + ) + if self.num_speculative_tokens > 0: + return await self._run_async_sched_step_overlap_mtp() + return await self._run_async_sched_step_overlap() @torch.inference_mode() def generate_output_tokens_dynamic_batch( self, loop: Optional[asyncio.AbstractEventLoop] = None ) -> Optional[Dict]: - """Synchronous wrapper for `self.async_generate_output_tokens_dynamic_batch.""" + """Synchronously run dynamic batching through any primer-only calls. + + Args: + loop (Optional[asyncio.AbstractEventLoop]): Event loop used to run + the asynchronous controller. + + Returns: + Optional[Dict]: Step output, or `None` when no work is active. + """ loop = get_asyncio_loop(loop) - return loop.run_until_complete(self.async_generate_output_tokens_dynamic_batch()) + while True: + context = self.inference_wrapped_model.inference_context + result = loop.run_until_complete( + self.async_generate_output_tokens_dynamic_batch( + run_async_overlap=context.can_prepare_requests() + ) + ) + if not result.primer_only: + return result.output def _update_top_n_logprobs_dict( self, @@ -2587,11 +3781,11 @@ def generate_all_output_tokens_static_batch( request.status = Status.COMPLETED + # Detokenize up to input_prompt_length + required_sequence_length for this idx. + sequence_length = input_prompt_length + required_sequence_length text, segments = self.detokenize_generations( - batch_prompt_tokens_with_generations[ - idx, : (input_prompt_length + required_sequence_length) - ], - input_prompt_length + generated_sequence_lengths, + batch_prompt_tokens_with_generations[idx, :sequence_length], + torch.tensor([sequence_length], device=batch_prompt_tokens_with_generations.device), sampling_params.return_segments, ) request.text = text # Inference server returns prompts & generations together diff --git a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/__init__.py b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/__init__.py index f2b0661dace..251039fa60a 100644 --- a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/__init__.py +++ b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/__init__.py @@ -5,7 +5,8 @@ from .chat_completions import bp as ChatCompletions from .completions import bp as Completions from .health import bp as Health + from .profile import bp as Profile - __all__ = [Completions, ChatCompletions, Health] + __all__ = [Completions, ChatCompletions, Health, Profile] except ImportError: __all__ = [] diff --git a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/chat_completions.py b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/chat_completions.py index 460acf39e9b..6f684a70dac 100644 --- a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/chat_completions.py +++ b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/chat_completions.py @@ -317,22 +317,6 @@ def _sanitize_tools_for_template(tools): return sanitized -def _reconstruct_reasoning_content(messages: list[dict]) -> list[dict]: - """Reconstruct tags from reasoning_content fields on assistant messages. - - For parity with vLLM, assistant messages may carry reasoning in the reasoning_content field. - Before applying the chat template, we must inline those tags back into content. - """ - for message in messages: - if message.get("role") != "assistant": - continue - reasoning_content = message.pop("reasoning_content", None) - if reasoning_content is not None: - content = message.get("content") or "" - message["content"] = f"{reasoning_content}{content}" - return messages - - def _replace_prefix_tokens( eos_token_id, previous_turn_token_ids, @@ -402,7 +386,9 @@ def _coerce_to_token_id_list(result): bp = Blueprint('chat_completions_api', __name__) - def apply_parsers(message_text, tools, parsers_list, tools_requested): + def apply_parsers( + message_text, tools, parsers_list, tools_requested, chat_template_kwargs=None + ): """Runs CPU-intensive text parsing.""" meta = {} for parser in parsers_list: @@ -410,7 +396,9 @@ def apply_parsers(message_text, tools, parsers_list, tools_requested): raise ValueError(f"Parser {parser} not found in PARSER_MAPPING") prev_text = message_text - parsed_text, new_info = PARSER_MAPPING[parser].parse(message_text, tools=tools) + parsed_text, new_info = PARSER_MAPPING[parser].parse( + message_text, tools=tools, chat_template_kwargs=chat_template_kwargs + ) if "tool_calls" in new_info: new_info["tool_calls"] = _normalize_tool_calls( new_info.get("tool_calls", []), tools=tools @@ -454,7 +442,6 @@ async def chat_completions(): if not isinstance(messages, list): return Response("'messages' must be a list", status=400) template_messages = _sanitize_messages_for_template(messages) - template_messages = _reconstruct_reasoning_content(template_messages) template_tools = _sanitize_tools_for_template(tools) try: @@ -584,6 +571,18 @@ async def chat_completions(): prompt_tokens = [tokenizer.bos] + prompt_tokens max_tokens = req.get("max_completion_tokens", None) or req.get("max_tokens", None) + ignore_eos = bool(req.get("ignore_eos", False)) + + # Does the client want the prompt tokens echoed back? Only then does the + # engine need to keep the prompt_tokens tensor on the response payload. + # return_tokenized_data (implied by prevent_retokenization) needs the ids; + # return_raw_text needs the ids to detokenize the prompt into raw_text. + prevent_retokenization = req.get("prevent_retokenization", True) + return_tokenized_data = ( + req.get("return_tokenized_data", False) or prevent_retokenization + ) + return_raw_text = req.get("return_raw_text", False) + return_prompt_tokens = return_tokenized_data or return_raw_text sampling_params = SamplingParams( temperature=temperature, @@ -594,6 +593,8 @@ async def chat_completions(): num_tokens_to_generate=(int(max_tokens) if max_tokens is not None else None), skip_prompt_log_probs=skip_prompt_log_probs, add_BOS=add_BOS, + termination_id=-1 if ignore_eos else None, + return_prompt_tokens=return_prompt_tokens, ) except ValueError as e: return Response(f"Invalid sampling parameter: {e}", status=400) @@ -654,15 +655,24 @@ async def chat_completions(): choices = [] total_completion_tokens = 0 prompt_tokens_counts = [] + cached_tokens_counts = [] + # return_tokenized_data / return_raw_text / return_prompt_tokens were computed + # at submit time (above) and drive both the response shape here and whether the + # engine kept the prompt_tokens tensor on the payload. request_idx = 0 for result_item in batch_results: result = unwrap_serialized_tensors(result_item) - prompt_tokens_out = result["prompt_tokens"] # The engine can modify prompt_tokens. text_output = result["generated_text"] - prompt_tokens_count = len(prompt_tokens_out) if prompt_tokens_out is not None else 0 + # The engine always reports prompt_length (for usage), but drops the + # prompt_tokens tensor unless return_prompt_tokens was set. + prompt_tokens_count = result.get("prompt_length") + if prompt_tokens_count is None: + prompt_tokens_out = result["prompt_tokens"] + prompt_tokens_count = len(prompt_tokens_out) if prompt_tokens_out is not None else 0 prompt_tokens_counts.append(prompt_tokens_count) + cached_tokens_counts.append(result.get("num_cached_tokens", 0)) logprobs_content = None if sampling_params.return_log_probs: @@ -703,7 +713,11 @@ async def chat_completions(): if parsers: message_text, metadata = apply_parsers( - message_text, tools, parsers, tools_requested + message_text, + tools, + parsers, + tools_requested, + chat_template_kwargs=chat_template_kwargs, ) normalized_tool_calls = metadata.get("tool_calls", []) @@ -728,9 +742,14 @@ async def chat_completions(): if "reasoning" in metadata: message["reasoning_content"] = metadata["reasoning"] - # Replicate data in the message field for compatibility. - message["prompt_token_ids"] = result["prompt_tokens"] - message["generation_token_ids"] = result["generated_tokens"] + if return_tokenized_data: + message["prompt_token_ids"] = result["prompt_tokens"] + message["generation_token_ids"] = result["generated_tokens"] + if return_raw_text: + prompt_str = tokenizer.detokenize(result["prompt_tokens"]) + message["raw_text"] = prompt_str + text_output + # Small RL/debug scalars (a few bytes each); harmless to keep for + # NeMo-RL compatibility. message["generation_log_probs"] = result.get("generated_log_probs", []) message["policy_epoch"] = result["policy_epoch"] message["kv_cache_epoch"] = result["kv_cache_epoch"] @@ -751,15 +770,13 @@ async def chat_completions(): else: finish_reason = "stop" + # Choice-level prompt/generation_token_ids, generation_log_probs and + # raw_text were duplicates of message-level data (or reconstructable); + # dropped to match vLLM's response shape and cut payload size. choice_data = { "index": request_idx, "message": message, - "prompt_token_ids": result["prompt_tokens"], - "generation_token_ids": result["generated_tokens"], - "generation_log_probs": result.get("generated_log_probs", []), - "raw_text": result["prompt"] + result["generated_text"], # 'logprobs' in chat API is an object containing 'content' - # "logprobs": {"content": logprobs_content} if logprobs_content else None, "logprobs": {"content": logprobs_content} if return_log_probs else None, "finish_reason": finish_reason, } @@ -774,7 +791,7 @@ async def chat_completions(): ] choices.append(choice_data) - if choice_data["generation_log_probs"] is None: + if result.get("generated_log_probs") is None: logger.warning( "Generation log probs is None for request:\n%s", json.dumps(_redact_token_id_lists_for_logging(result), indent=4), @@ -783,6 +800,7 @@ async def chat_completions(): request_idx += 1 prompt_token_count = max(prompt_tokens_counts) if prompt_tokens_counts else 0 + cached_token_count = max(cached_tokens_counts) if cached_tokens_counts else 0 response = { "id": f"chatcmpl-{uuid.uuid4().hex}", "created": int(time.time()), @@ -793,6 +811,7 @@ async def chat_completions(): "prompt_tokens": prompt_token_count, "completion_tokens": total_completion_tokens, "total_tokens": prompt_token_count + total_completion_tokens, + "prompt_tokens_details": {"cached_tokens": cached_token_count}, }, } diff --git a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/completions.py b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/completions.py index 6f57a863c1c..2e2a57d6fc1 100644 --- a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/completions.py +++ b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/completions.py @@ -122,6 +122,9 @@ async def completions(): num_tokens_to_generate=sampling_params.num_tokens_to_generate, stop_words=sampling_params.stop_words, termination_id=sampling_params.termination_id, + # This endpoint always echoes prompt_token_ids in its response, so + # keep the prompt tokens on the payload (default is now to drop them). + return_prompt_tokens=True, ) tasks.append(client.add_request(prompt_tokens, per_req_params)) diff --git a/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/profile.py b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/profile.py new file mode 100644 index 00000000000..71478aace70 --- /dev/null +++ b/megatron/core/inference/text_generation_server/dynamic_text_gen_server/endpoints/profile.py @@ -0,0 +1,41 @@ +# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. + +"""CUDA profiler control endpoints. + +POST /start_profile and /stop_profile relay a control signal through the +InferenceClient -> data-parallel coordinator -> every connected EP/DP engine, +which calls cudaProfilerStart()/cudaProfilerStop(). Pair with an outer +`nsys profile --capture-range=cudaProfilerApi` to bracket a capture window. +""" + +import logging + +logger = logging.getLogger(__name__) + +try: + from quart import Blueprint, current_app, jsonify + + bp = Blueprint('profile_api', __name__) + + @bp.route('/start_profile', methods=['POST']) + @bp.route('/v1/start_profile', methods=['POST']) + async def start_profile(): + """Broadcast cudaProfilerStart to all engines.""" + client = current_app.config.get('client') + if client is None: + return jsonify({"status": "error", "details": "client not initialized"}), 503 + client.start_cuda_profiler() + return jsonify({"status": "ok", "action": "start_profile"}), 200 + + @bp.route('/stop_profile', methods=['POST']) + @bp.route('/v1/stop_profile', methods=['POST']) + async def stop_profile(): + """Broadcast cudaProfilerStop to all engines.""" + client = current_app.config.get('client') + if client is None: + return jsonify({"status": "error", "details": "client not initialized"}), 503 + client.stop_cuda_profiler() + return jsonify({"status": "ok", "action": "stop_profile"}), 200 + +except ImportError as e: + logger.warning(f"Could not import quart: {e}") diff --git a/megatron/core/model_parallel_config.py b/megatron/core/model_parallel_config.py index 8a943b3ef2b..85acd260b58 100644 --- a/megatron/core/model_parallel_config.py +++ b/megatron/core/model_parallel_config.py @@ -24,6 +24,41 @@ def _parse_pad_packed_seq_alignment(value): ) from exc +def resolve_tensor_parallel_weight_shards( + tensor_model_parallel_size: int, + tensor_parallel_num_weight_shards: Optional[int], + gtp_weight_remat_size: int, + shards_field: str = "tensor_parallel_num_weight_shards", + tp_field: str = "tensor_model_parallel_size", +) -> tuple: + """Reconcile ``tensor_parallel_num_weight_shards`` and ``gtp_weight_remat_size``. + + ``tensor_parallel_num_weight_shards`` is the user-facing total number of shards each weight is + split into across the tensor-parallel + GTP axes. It is the source of truth and implies + ``gtp_weight_remat_size = tensor_parallel_num_weight_shards // tensor_model_parallel_size``. + When None it defaults to ``tensor_model_parallel_size * gtp_weight_remat_size`` (so the pair + stays consistent, and equals ``tensor_model_parallel_size`` in the no-GTP default). Idempotent. + + Returns the reconciled ``(tensor_parallel_num_weight_shards, gtp_weight_remat_size)``. + """ + tp = tensor_model_parallel_size + if tensor_parallel_num_weight_shards is None: + tensor_parallel_num_weight_shards = tp * gtp_weight_remat_size + else: + if tensor_parallel_num_weight_shards < tp: + raise ValueError( + f"{shards_field} ({tensor_parallel_num_weight_shards}) must be " + f">= {tp_field} ({tp})." + ) + if tensor_parallel_num_weight_shards % tp != 0: + raise ValueError( + f"{shards_field} ({tensor_parallel_num_weight_shards}) must be " + f"divisible by {tp_field} ({tp})." + ) + gtp_weight_remat_size = tensor_parallel_num_weight_shards // tp + return tensor_parallel_num_weight_shards, gtp_weight_remat_size + + @dataclass @experimental_api class ModelParallelConfig: @@ -38,6 +73,26 @@ class ModelParallelConfig: tensor_model_parallel_size: int = 1 """Intra-layer model parallelism. Splits tensors across GPU ranks.""" + tensor_parallel_num_weight_shards: Optional[int] = None + """Total number of shards each weight is split into across the tensor-parallel + GTP axes + (i.e. ``tensor_model_parallel_size * gtp_weight_remat_size``). This is the user-facing knob: + it must be ``>= tensor_model_parallel_size`` and divisible by it. When None it defaults to + ``tensor_model_parallel_size`` (no GTP sharding). It is the source of truth and implies + ``gtp_weight_remat_size = tensor_parallel_num_weight_shards // tensor_model_parallel_size`` + (resolved in ``__post_init__``). + """ + + gtp_weight_remat_size: int = 1 + """Generalized tensor parallelism with weight rematerialization. Shards model weights + across GPU ranks along ``out_features``; each weight is rematerialized independently + (per-weight, not per-layer) via async all-gather on every forward AND backward pass. + Placed right after tensor parallelism in the parallelism ordering. + + INTERNAL / DERIVED — there is no CLI flag for it; do not set directly. It is computed in + ``__post_init__`` from ``tensor_parallel_num_weight_shards`` (= that value divided by + ``tensor_model_parallel_size``). Use ``tensor_parallel_num_weight_shards`` to control GTP. + """ + pipeline_model_parallel_comm_backend: Optional[Literal["nccl", "ucc"]] = None """Configuring backend option of pipeline parallel communication (e.g., nccl, ucc) If None, the default backend will be used. @@ -143,6 +198,27 @@ class ModelParallelConfig: Default is None, which will be set to the value of tensor_model_parallel_size. """ + expert_tensor_parallel_num_weight_shards: Optional[int] = None + """Total number of shards each expert weight is split into across the expert-tensor-parallel + + expert-GTP axes (i.e. ``expert_tensor_parallel_size * expert_gtp_weight_remat_size``). This + is the user-facing knob for expert layers: it must be ``>= expert_tensor_parallel_size`` and + divisible by it. When None it defaults to ``expert_tensor_parallel_size`` (no expert GTP + sharding). It is the source of truth and implies + ``expert_gtp_weight_remat_size = expert_tensor_parallel_num_weight_shards // + expert_tensor_parallel_size`` (resolved in ``__post_init__``). + """ + + expert_gtp_weight_remat_size: int = 1 + """Generalized tensor parallelism with weight rematerialization, for expert layers. Independent + from the decoder's ``gtp_weight_remat_size``. + Placed right after expert parallelism in the parallelism ordering. + + INTERNAL / DERIVED — there is no CLI flag for it; do not set directly. It is computed in + ``__post_init__`` from ``expert_tensor_parallel_num_weight_shards`` (= that value divided by + ``expert_tensor_parallel_size``). Use ``expert_tensor_parallel_num_weight_shards`` to control + expert GTP. + """ + ################### # Initialization ################### @@ -557,6 +633,24 @@ def __post_init__(self): if self.expert_tensor_parallel_size is None: self.expert_tensor_parallel_size = self.tensor_model_parallel_size + # Derive the internal gtp_weight_remat_size from the user-facing + # tensor_parallel_num_weight_shards: + # num_weight_shards = tensor_model_parallel_size * gtp_weight_remat + _, self.gtp_weight_remat_size = resolve_tensor_parallel_weight_shards( + self.tensor_model_parallel_size, + self.tensor_parallel_num_weight_shards, + self.gtp_weight_remat_size, + ) + + # Same reconciliation for expert layers (expert_tensor_parallel_size finalized above). + _, self.expert_gtp_weight_remat_size = resolve_tensor_parallel_weight_shards( + self.expert_tensor_parallel_size, + self.expert_tensor_parallel_num_weight_shards, + self.expert_gtp_weight_remat_size, + shards_field="expert_tensor_parallel_num_weight_shards", + tp_field="expert_tensor_parallel_size", + ) + if self.pipeline_model_parallel_size > 1: if self.pipeline_dtype is None: raise ValueError( diff --git a/megatron/core/models/audio/__init__.py b/megatron/core/models/audio/__init__.py index f0cd2776ff9..538e117fe98 100644 --- a/megatron/core/models/audio/__init__.py +++ b/megatron/core/models/audio/__init__.py @@ -5,6 +5,7 @@ NemoTransformerAudioTokenEstimator, ceil_div, ) +from .audio_processor import NemoAudioProcessor from .audio_projector import AudioProjection from .nemo_audio_checkpoint import ( CHECKPOINT_NEMO_AUDIO_PREPROCESSOR_CONFIG_NAME, @@ -29,6 +30,7 @@ "CHECKPOINT_NEMO_AUDIO_PREPROCESSOR_CONFIG_NAME", "CHECKPOINT_NEMO_TRANSFORMER_AUDIO_CONFIG_NAME", "NemoAudioFeatureConfig", + "NemoAudioProcessor", "NemoTransformerAudioConfig", "NemoTransformerAudioModel", "NemoTransformerAudioTokenEstimator", diff --git a/megatron/core/models/audio/audio_processor.py b/megatron/core/models/audio/audio_processor.py new file mode 100644 index 00000000000..30ccaabd257 --- /dev/null +++ b/megatron/core/models/audio/audio_processor.py @@ -0,0 +1,338 @@ +# Copyright (c) 2025, NVIDIA CORPORATION. +# SPDX-License-Identifier: BSD-3-Clause + +"""Waveform audio processor for the NeMo Transformer audio frontend. + +``NemoAudioProcessor`` is the concrete, data-side feature extractor that the +multimodal data pipeline uses to (a) estimate how many encoder/projector +embeddings an audio clip expands to (for placeholder expansion and packing) and +(b) materialize log-mel features for the audio encoder. It composes the +model-frontend descriptors (``NemoAudioFeatureConfig`` and +``NemoTransformerAudioTokenEstimator``) with the vendored standalone log-mel +preprocessor. + +The data pipeline depends on this only through a small duck-typed interface +(``compute_num_embeddings`` / ``compute_num_frames`` / ``materialize`` plus the +cumulative-prefix ``num_*_from_num_samples`` primitives), so an ``audio_ref`` is +treated structurally — this module has no dependency on the data library's +``AudioRef`` type. +""" + +from __future__ import annotations + +from typing import Any + +import torch +import torch.nn.functional as F + +from megatron.core.models.audio.audio_feature_config import ( + NemoAudioFeatureConfig, + NemoTransformerAudioTokenEstimator, +) + +_AUDIO_DURATION_MISMATCH_TOLERANCE_SECONDS = 0.5 + + +def _load_waveform_from_spec(audio_spec: dict[str, Any]) -> tuple[torch.Tensor, int | None]: + kind = audio_spec.get("kind") + if kind == "avdecoder": + return _decode_avdecoder( + audio_spec["decoder"], + audio_spec.get("source_name", ""), + sample_rate=audio_spec.get("sample_rate") or audio_spec.get("sampling_rate"), + ) + raise ValueError(f"Unsupported audio kind {kind!r}") + + +def _resolve_lazy_media(media: Any) -> Any: + if hasattr(media, "get") and not hasattr(media, "get_audio"): + media = media.get() + if isinstance(media, (list, tuple)): + if not media: + raise ValueError("Lazy audio media resolved to an empty sequence.") + media = media[0] + return media + + +def _audio_clip_to_float32(clip: torch.Tensor) -> torch.Tensor: + if not torch.is_tensor(clip): + clip = torch.as_tensor(clip) + if clip.ndim == 1: + clip = clip.unsqueeze(0) + elif clip.ndim != 2: + raise ValueError(f"Unsupported decoded audio clip shape {tuple(clip.shape)}.") + + if clip.dtype.is_floating_point: + return clip.to(torch.float32).contiguous() + + if clip.dtype == torch.uint8: + return ((clip.to(torch.float32) - 128.0) / 128.0).contiguous() + + if clip.dtype in (torch.int8, torch.int16, torch.int32, torch.int64): + scale = float(torch.iinfo(clip.dtype).max) + return (clip.to(torch.float32) / scale).contiguous() + + raise ValueError(f"Unsupported decoded audio dtype {clip.dtype}.") + + +def _decoder_sample_rate(decoder: Any, sample_rate: int | None) -> int | None: + if sample_rate is not None: + return int(sample_rate) + + if hasattr(decoder, "get_audio_samples_per_second"): + return int(decoder.get_audio_samples_per_second()) + + if hasattr(decoder, "get_metadata"): + metadata = decoder.get_metadata( + get_video=False, + get_video_duration=False, + get_video_frame_count=False, + get_video_frame_size=False, + get_audio=True, + get_audio_duration=False, + ) + audio_sample_rate = getattr(metadata, "audio_sample_rate", None) + if audio_sample_rate is not None: + return int(audio_sample_rate) + + return None + + +def _decode_avdecoder( + decoder: Any, source_name: str, *, sample_rate: int | None = None +) -> tuple[torch.Tensor, int | None]: + decoder = _resolve_lazy_media(decoder) + if not hasattr(decoder, "get_audio"): + raise ValueError( + f"Expected AVDecoder-like audio media for {source_name}, " + f"got {type(decoder).__name__}." + ) + + av_data = decoder.get_audio() + clips = getattr(av_data, "audio_clips", None) + if not clips: + raise ValueError(f"Decoded audio {source_name!r} did not contain audio clips.") + + waveform = torch.cat([_audio_clip_to_float32(clip) for clip in clips], dim=-1) + return waveform.contiguous(), _decoder_sample_rate(decoder, sample_rate) + + +def _resolve_sample_rate(audio_ref: Any, decoded_sample_rate: int | None) -> int | None: + sample_rate = audio_ref.sample_rate + if sample_rate is None: + sample_rate = decoded_sample_rate + if sample_rate is None and isinstance(audio_ref.data, dict): + sample_rate = audio_ref.data.get("sample_rate") or audio_ref.data.get("sampling_rate") + if sample_rate is None: + return None + return int(sample_rate) + + +def _audio_num_sample_tolerance(audio_ref: Any, decoded_sample_rate: int | None) -> int: + sample_rate = _resolve_sample_rate(audio_ref, decoded_sample_rate) + if sample_rate is None: + return 0 + return int(_AUDIO_DURATION_MISMATCH_TOLERANCE_SECONDS * sample_rate + 0.999999) + + +def _normalize_mono_waveform(audio_ref: Any) -> tuple[torch.Tensor, int | None]: + data = audio_ref.data + decoded_sample_rate = None + if torch.is_tensor(data): + waveform = data + elif isinstance(data, dict): + waveform, decoded_sample_rate = _load_waveform_from_spec(data) + else: + raise ValueError( + "AudioRef.data must be a raw float32 waveform tensor or a supported lazy " + "audio spec for the Megatron multimodal audio path." + ) + if waveform.dtype != torch.float32: + raise ValueError(f"Expected raw float32 waveform tensor, got {waveform.dtype}.") + + if waveform.ndim == 1: + pass + elif waveform.ndim == 2: + waveform = waveform.mean(dim=0) + else: + raise ValueError( + f"Unsupported waveform shape {tuple(waveform.shape)}. Expected [T] or [C, T]." + ) + + # First reconcile the decoded waveform to the full source length (num_samples + # always counts the un-sliced source), then crop to slice_range if set. + if audio_ref.num_samples is not None: + num_samples = int(audio_ref.num_samples) + available_num_samples = int(waveform.shape[0]) + diff = num_samples - available_num_samples + tolerance = _audio_num_sample_tolerance(audio_ref, decoded_sample_rate) + if diff > tolerance: + raise ValueError( + f"AudioRef.num_samples={num_samples} exceeds waveform length " + f"{available_num_samples} by {diff} samples, which is greater than " + f"the allowed tolerance {tolerance}." + ) + if diff > 0: + waveform = F.pad(waveform, (0, diff)) + elif diff < 0: + waveform = waveform[:num_samples] + + if audio_ref.slice_range is not None: + start, end = int(audio_ref.slice_range[0]), int(audio_ref.slice_range[1]) + if start < 0 or end < start: + raise ValueError( + f"AudioRef.slice_range must satisfy 0 <= start <= end, got {(start, end)}" + ) + waveform = waveform[start:end] + + return waveform.contiguous(), decoded_sample_rate + + +def _infer_num_samples(audio_ref: Any) -> int: + # slice_range, when set, defines the effective length of this ref. + if audio_ref.slice_range is not None: + start, end = int(audio_ref.slice_range[0]), int(audio_ref.slice_range[1]) + return max(0, end - start) + if audio_ref.num_samples is not None: + return int(audio_ref.num_samples) + + data = audio_ref.data + if torch.is_tensor(data): + waveform = data + elif isinstance(data, dict): + waveform, _ = _load_waveform_from_spec(data) + else: + raise ValueError( + "AudioRef.data must be a raw float32 waveform tensor or a supported lazy " + "audio spec for the Megatron multimodal audio path." + ) + if waveform.dtype != torch.float32: + raise ValueError(f"Expected raw float32 waveform tensor, got {waveform.dtype}.") + + if waveform.ndim == 1: + available_num_samples = int(waveform.shape[0]) + elif waveform.ndim == 2: + available_num_samples = int(waveform.shape[-1]) + else: + raise ValueError( + f"Unsupported waveform shape {tuple(waveform.shape)}. Expected [T] or [C, T]." + ) + + return available_num_samples + + +class NemoAudioProcessor: + """Waveform audio processor with a NeMo log-mel frontend.""" + + def __init__( + self, + *, + token_estimator: NemoTransformerAudioTokenEstimator, + feature_config: NemoAudioFeatureConfig | None = None, + ) -> None: + # Lazy import keeps construction cheap and avoids importing the heavier + # preprocessor module until a processor is actually built. The vendored + # ``AudioToMelSpectrogramPreprocessor`` is a stdlib+PyTorch port of NeMo's + # preprocessor; ``.eval()`` disables training-time dither and narrowband + # augmentation (typical for a feature extractor inside the multimodal + # pipeline; flip back via ``.train()`` if needed). + from megatron.core.models.audio.nemo_audio_preprocessing import ( + AudioToMelSpectrogramPreprocessor, + ) + + self.token_estimator = token_estimator + self.feature_config = feature_config or NemoAudioFeatureConfig() + self._preprocessor = AudioToMelSpectrogramPreprocessor( + **self.feature_config.to_nemo_kwargs() + ).eval() + # The vendored standalone AudioToMelSpectrogramPreprocessor exposes + # win/hop lengths directly (no ``featurizer`` indirection). + self._hop_length = int(self._preprocessor.hop_length) + self._n_mels = int(self.feature_config.features) + self._sample_rate = int(self.feature_config.sample_rate) + + @property + def input_feature_dim(self) -> int: + """Number of mel feature bins produced per frame (the encoder input dim).""" + return self._n_mels + + @property + def sample_rate(self) -> int: + """Expected input waveform sample rate, in Hz.""" + return self._sample_rate + + def _validate_sample_rate(self, audio_ref: Any, decoded_sample_rate: int | None = None) -> None: + sample_rate = audio_ref.sample_rate + if sample_rate is None: + sample_rate = decoded_sample_rate + if sample_rate is not None and int(sample_rate) != self.sample_rate: + raise ValueError( + f"Expected audio sample rate {self.sample_rate}, got {sample_rate}. " + "Resample raw waveforms to the encoder sample rate before packing." + ) + + def _compute_num_frames_from_num_samples(self, num_samples: int) -> int: + if num_samples < 0: + raise ValueError(f"num_samples must be >= 0, got {num_samples}") + if num_samples == 0: + return 0 + return int(num_samples // self._hop_length) + + def compute_num_frames(self, audio_ref: Any) -> int: + """Number of feature frames the clip described by ``audio_ref`` expands to.""" + self._validate_sample_rate(audio_ref) + return self._compute_num_frames_from_num_samples(_infer_num_samples(audio_ref)) + + def num_frames_from_num_samples(self, num_samples: int) -> int: + """Pure frame-count math for an audio prefix of ``num_samples`` samples. + + Slice math: + ``frames_in([s, e)) = num_frames_from_num_samples(e) - num_frames_from_num_samples(s)``. + """ + return self._compute_num_frames_from_num_samples(num_samples) + + def num_embeddings_from_num_samples(self, num_samples: int) -> int: + """Pure embedding-count math for an audio prefix of ``num_samples`` samples. + + Slice math: ``embeds_in([s, e))`` = + ``num_embeddings_from_num_samples(e) - num_embeddings_from_num_samples(s)``. + """ + return self.token_estimator.estimate_from_num_frames( + self._compute_num_frames_from_num_samples(num_samples) + ) + + def compute_num_embeddings(self, audio_ref: Any) -> int: + """Number of encoder/projector embeddings the clip in ``audio_ref`` expands to.""" + self._validate_sample_rate(audio_ref) + return self.token_estimator.estimate_from_num_frames(self.compute_num_frames(audio_ref)) + + def materialize(self, audio_ref: Any) -> tuple[torch.Tensor, int]: + """Decode ``audio_ref`` and return its ``(T, n_mels)`` log-mel features and frame count.""" + waveform, decoded_sample_rate = _normalize_mono_waveform(audio_ref) + self._validate_sample_rate(audio_ref, decoded_sample_rate) + + num_samples = waveform.shape[0] + num_frames = self._compute_num_frames_from_num_samples(num_samples) + if num_frames == 0: + return ( + torch.empty( + (0, self.input_feature_dim), dtype=torch.float32, device=waveform.device + ), + 0, + ) + + batched = waveform.unsqueeze(0) + lengths = torch.tensor([num_samples], dtype=torch.long, device=waveform.device) + mels, out_lengths = self._preprocessor(batched, lengths) + + # mels: (1, n_mels, T_frames) -> (T_frames, n_mels), trimmed to valid frames. + valid_frames = int(out_lengths[0].item()) + if valid_frames == 0: + return ( + torch.empty( + (0, self.input_feature_dim), dtype=torch.float32, device=waveform.device + ), + 0, + ) + log_mel = mels[0, :, :valid_frames].transpose(0, 1).contiguous() + return log_mel.to(torch.float32), valid_frames diff --git a/megatron/core/models/backends.py b/megatron/core/models/backends.py index a270161ddd6..e29265c04f4 100644 --- a/megatron/core/models/backends.py +++ b/megatron/core/models/backends.py @@ -4,7 +4,7 @@ import warnings from abc import abstractmethod from functools import partial -from typing import Optional, Protocol, cast +from typing import Literal, Optional, Protocol, cast from megatron.core.extensions.transformer_engine import ( TEColumnParallelGroupedLinear, @@ -200,3 +200,19 @@ def grouped_mlp_modules(self, moe_use_grouped_gemm: bool) -> ExpertsBuilder: activation_func=self.activation_func(), ), ) + + +def get_backend( + transformer_impl: Literal["local", "transformer_engine", "inference_optimized"] +) -> BackendSpecProvider: + """Return the backend that's selected with the given `transformer_impl`.""" + if transformer_impl == "transformer_engine": + from megatron.core.extensions.transformer_engine_spec_provider import TESpecProvider + + return TESpecProvider() + elif transformer_impl == "inference_optimized": + return InferenceSpecProvider() + elif transformer_impl == "local": + return LocalSpecProvider() + else: + raise ValueError(f"unknown transformer_impl='{transformer_impl}'") diff --git a/megatron/core/models/bert/bert_model.py b/megatron/core/models/bert/bert_model.py index 3fd1e01f4a1..9260697b600 100644 --- a/megatron/core/models/bert/bert_model.py +++ b/megatron/core/models/bert/bert_model.py @@ -49,6 +49,13 @@ class BertModel(LanguageModule): rotary_percent (float): Percent of rotary dimension to use for rotary position embeddings. Defaults to 1.0 (100%). Ignored unless position_embedding_type is 'rope'. vp_stage (int): Virtual pipeline stage. + apply_lm_head (bool): Whether to transform the encoder's final hidden states with + ``BertLMHead`` (dense + GeLU + LayerNorm) before the vocabulary projection. + Defaults to True. Set to False for architectures whose output projection is + applied directly to the encoder output (e.g. models with their own final norm), + bypassing BERT's dense+GeLU+LayerNorm transform. + output_layer_bias (bool): Whether to include a bias in the vocabulary projection. + Defaults to True for backward compatibility. """ def __init__( @@ -70,6 +77,8 @@ def __init__( return_embeddings=False, vp_stage: Optional[int] = None, pg_collection: Optional[ProcessGroupCollection] = None, + apply_lm_head: bool = True, + output_layer_bias: bool = True, ): super(BertModel, self).__init__(config=config, pg_collection=pg_collection) @@ -92,6 +101,8 @@ def __init__( self.add_binary_head = add_binary_head self.return_embeddings = return_embeddings self.vp_stage = vp_stage + self.apply_lm_head = apply_lm_head + self.output_layer_bias = output_layer_bias # megatron core pipelining currently depends on model type self.model_type = ModelType.encoder_or_decoder @@ -129,7 +140,7 @@ def __init__( # Output if post_process: # TODO: Make sure you are passing in the mpu_vocab_size properly - self.lm_head = BertLMHead(config.hidden_size, config) + self.lm_head = BertLMHead(config.hidden_size, config) if self.apply_lm_head else None self.output_layer = tensor_parallel.ColumnParallelLinear( config.hidden_size, @@ -140,7 +151,7 @@ def __init__( if config.use_mup and not self.share_embeddings_and_output_weights else config.init_method ), - bias=True, + bias=self.output_layer_bias, skip_bias_add=False, gather_output=not self.parallel_output, skip_weight_param_allocation=pre_process and share_embeddings_and_output_weights, @@ -375,7 +386,9 @@ def forward( if self.share_embeddings_and_output_weights: output_weight = self.shared_embedding_or_output_weight() - hidden_states_after_lm_head = self.lm_head(hidden_states=hidden_states) + hidden_states_after_lm_head = ( + self.lm_head(hidden_states=hidden_states) if self.lm_head is not None else hidden_states + ) logits, _ = self.output_layer(hidden_states_after_lm_head, weight=output_weight) binary_logits = None diff --git a/megatron/core/models/common/fine_grained_callables.py b/megatron/core/models/common/fine_grained_callables.py new file mode 100644 index 00000000000..8f46711d553 --- /dev/null +++ b/megatron/core/models/common/fine_grained_callables.py @@ -0,0 +1,169 @@ +# Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Layer-callable builders for the combined-1F1B fine-grained schedule plan. + +These build_* functions assemble the per-layer ``(forward_funcs, backward_dw)`` +tuple that the schedule plan plugs into ``TransformerLayerNode``. + +The TransformerLayer-specific builder lives in ``gpt/fine_grained_callables.py`` +because it depends on GPT's MoE wiring; the MTP builder and the dispatcher +``build_layer_callables`` are model-agnostic — both GPTModel and HybridModel +schedule MTP layers identically — so they live here. +""" + +from contextlib import nullcontext +from functools import partial + +import torch + +from megatron.core import tensor_parallel +from megatron.core.models.gpt.fine_grained_callables import build_transformer_layer_callables +from megatron.core.transformer.moe.moe_layer import MoELayer +from megatron.core.transformer.multi_token_prediction import ( + MultiTokenPredictionLayer, + get_mtp_layer_offset, +) +from megatron.core.transformer.transformer_layer import TransformerLayer, make_viewless_tensor + + +def build_mtp_layer_callables(layer): + """Callables for multi-token prediction layer nodes. + + Wraps the inner ``layer.mtp_model_layer``'s callables with MTP-specific + pre-process (chunk and concat embeddings) and post-process (gather across + depths) steps. + """ + + forward_funcs, backward_dw = build_layer_callables(layer.mtp_model_layer) + is_moe, _ = get_layer_moe_metadata(layer.mtp_model_layer) + (pre_dispatch_forward, dispatch_forward, mlp_forward, combine_forward, _) = forward_funcs + assert is_moe, "MTP layer in a2a overlap only supports MoE layer for now." + + def submodule_mtp_pre_dispatch_forward(node, hidden_states): + # MTP Block Preprocess + if node.is_first_layer: + # Apply the main decoder's final_norm if this VPP chunk owns it but + # holds no main HybridStack layers — without this, ``_maybe_apply_final_norm`` + # never fires for the main path and the unnormalized hidden_states feed + # straight into the LM head (lm_loss explodes by ~10x; grads diverge). + # Restricted to HybridModel because GPT models go through a different + # MTP wiring path. Must run before ``torch.chunk`` so every chunk — + # including the main-decoder slice consumed by the LM head — sees + # the norm; the MTP slices then go through MTP's own ``hnorm`` as usual. + from megatron.core.models.hybrid.hybrid_model import HybridModel + + model = node.chunk_state.model + if isinstance(model, HybridModel) and len(model.decoder.layers) == 0: + final_norm = getattr(model.decoder, "final_norm", None) or getattr( + model.decoder, "final_layernorm", None + ) + if final_norm is not None: + hidden_states = final_norm(hidden_states) + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) + + offset = get_mtp_layer_offset(layer.config, node.chunk_state.model.vp_stage) + node.chunk_state.mtp_hidden_states = list(torch.chunk(hidden_states, 1 + offset, dim=0)) + hidden_states = node.chunk_state.mtp_hidden_states[offset] + + input_ids, position_ids, padding_mask, decoder_input, hidden_states = layer._get_embeddings( + input_ids=node.chunk_state.input_ids, + position_ids=node.chunk_state.position_ids, + embedding=node.chunk_state.model.embedding, + hidden_states=hidden_states, + packed_seq_params=node.chunk_state.packed_seq_params, + padding_mask=node.chunk_state.padding_mask, + ) + node.chunk_state.input_ids = input_ids + node.chunk_state.position_ids = position_ids + node.chunk_state.padding_mask = padding_mask + + # MTP Layer Preprocess + # norm, linear projection and transformer + assert ( + node.chunk_state.context is None + ), f"multi token prediction + cross attention is not yet supported." + assert ( + node.chunk_state.packed_seq_params is None + ), f"multi token prediction + sequence packing is not yet supported." + + if layer.config.sequence_parallel: + rng_context = tensor_parallel.get_cuda_rng_tracker().fork() + else: + rng_context = nullcontext() + + # fp8 context is added in 1f1b schedule, so we don't need to add it here + with rng_context: + hidden_states = layer._concat_embeddings(hidden_states, decoder_input) + return pre_dispatch_forward(node, hidden_states) + + def submodule_mtp_postprocess_forward(node, hidden_states): + hidden_states = layer._postprocess(hidden_states) + node.chunk_state.mtp_hidden_states.append(hidden_states) + if node.is_last_layer: + hidden_states = torch.cat(node.chunk_state.mtp_hidden_states, dim=0) + node.chunk_state.mtp_hidden_states = None + return hidden_states + + def rng_context_wrapper(func, *args, **kwargs): + """ + Wrapper to add rng context to submodule callables + """ + if layer.config.sequence_parallel: + rng_context = tensor_parallel.get_cuda_rng_tracker().fork() + else: + rng_context = nullcontext() + with rng_context: + return func(*args, **kwargs) + + # Build forward and backward callable functions. + # pre_dispatch_func already has rng context (rolled into + # submodule_mtp_pre_dispatch_forward), so it does not need to be wrapped. + pre_dispatch_func = submodule_mtp_pre_dispatch_forward + dispatch_func = partial(rng_context_wrapper, dispatch_forward) + mlp_func = partial(rng_context_wrapper, mlp_forward) + combine_func = partial(rng_context_wrapper, combine_forward) + mtp_post_process_func = submodule_mtp_postprocess_forward + + forward_funcs = [ + pre_dispatch_func, + dispatch_func, + mlp_func, + combine_func, + mtp_post_process_func, + ] + pre_dispatch_bwd = backward_dw["pre_dispatch_computation"] + if isinstance(pre_dispatch_bwd, list): + pre_dispatch_bwd.append(layer.eh_proj) + else: + backward_dw["pre_dispatch_computation"] = [pre_dispatch_bwd, layer.eh_proj] + + return forward_funcs, backward_dw + + +def get_layer_moe_metadata(layer): + """Return ``(is_moe, num_local_experts)`` for schedule-node construction.""" + + if isinstance(layer, MultiTokenPredictionLayer): + return get_layer_moe_metadata(layer.mtp_model_layer) + if isinstance(layer, TransformerLayer): + is_moe = isinstance(layer.mlp, MoELayer) + num_local_experts = layer.mlp.num_local_experts if is_moe else None + return is_moe, num_local_experts + + raise ValueError(f"Unsupported layer type: {type(layer)}") + + +def build_layer_callables(layer): + """Dispatch to the appropriate layer-callable builder. + + Returns ``(forward_funcs, backward_dw)``. + """ + + if isinstance(layer, MultiTokenPredictionLayer): + return build_mtp_layer_callables(layer) + if isinstance(layer, TransformerLayer): + return build_transformer_layer_callables(layer) + + raise ValueError(f"Unsupported layer type: {type(layer)}") diff --git a/megatron/core/models/common/language_module/language_module.py b/megatron/core/models/common/language_module/language_module.py index 92db84ce0c9..ae70a502e75 100644 --- a/megatron/core/models/common/language_module/language_module.py +++ b/megatron/core/models/common/language_module/language_module.py @@ -141,6 +141,17 @@ def check_and_set_env_variable( check_and_set_env_variable("NVTE_FUSED_ATTN", 1, AttnBackend.auto) check_and_set_env_variable("NVTE_UNFUSED_ATTN", 1, AttnBackend.auto) + # Pin the FlashAttention generation for TransformerEngine by disabling the + # other versions via NVTE_FLASH_ATTN_V2/V3/V4 (default 1). This keeps the + # training-side attention on the same kernel as the mcore inference path, + # which honors config.flash_attention_version directly. + if self.config.flash_attention_version is not None: + for version in (2, 3, 4): + if version != self.config.flash_attention_version: + check_and_set_env_variable( + f"NVTE_FLASH_ATTN_V{version}", 0, self.config.attention_backend + ) + def compute_language_model_loss(self, labels: Tensor, logits: Tensor) -> Tensor: """Computes the language model loss (Cross entropy across vocabulary) diff --git a/megatron/core/models/common/model_chunk_schedule_plan.py b/megatron/core/models/common/model_chunk_schedule_plan.py index 45e7f66dfe6..0feb8577adf 100644 --- a/megatron/core/models/common/model_chunk_schedule_plan.py +++ b/megatron/core/models/common/model_chunk_schedule_plan.py @@ -35,7 +35,8 @@ class TransformerLayerSchedulePlan: MLP, MoE dispatch and combine, optional mHC recomputation, and MTP post-processing nodes. layer (TransformerLayerSchedulePlan) - ├── attn (TransformerLayerNode): attention -> layernorm -> router -> dispatch preprocess + ├── pre_dispatch_computation (TransformerLayerNode): + │ attention -> layernorm -> router -> dispatch preprocess ├── moe_dispatch (TransformerLayerNode): dispatch All2All ├── mlp (TransformerLayerNode): mlp module ├── moe_combine (TransformerLayerNode): combine All2All (incl. MLP-side mHC post-processing) @@ -43,13 +44,15 @@ class TransformerLayerSchedulePlan: └── mtp_post_process (PostProcessNode): mtp post process Note that MTP layer has the same operation and execution order with TransformerLayer regarding - moe_dispatch, mlp, moe_combine, but contains extra operations in attn and mtp_post_process: - * mtp.attn wraps around transformer_layer.attn with extra norm, proj and embedding operations. + moe_dispatch, mlp, moe_combine, but contains extra operations in + pre_dispatch_computation and mtp_post_process: + * mtp.pre_dispatch_computation wraps around transformer_layer.pre_dispatch_computation with + extra norm, proj and embedding operations. * mtp.mtp_post_process contains output_layer, mtp loss operations, whereas transformer_layer.mtp_post_process is empty. """ - attn = None + pre_dispatch_computation = None moe_dispatch = None mlp = None moe_combine = None @@ -72,10 +75,10 @@ def __init__(self, layer, event, chunk_state, comp_stream, comm_stream, extra_ar The event and chunk_state are binded to the TransformerModelChunkSchedulePlan and shared across all layers in the model chunk. """ - from megatron.core.models.gpt.fine_grained_callables import TransformerLayerState + from megatron.core.models.common.utils import LayerState self.config = layer.config - self.layer_state = TransformerLayerState() + self.layer_state = LayerState() self.chunk_state = chunk_state self.layer = layer self.event = event @@ -87,9 +90,9 @@ def __init__(self, layer, event, chunk_state, comp_stream, comm_stream, extra_ar def release_state(self): """Release reference, this helps avoid memory leak.""" - if hasattr(self, 'attn') and self.attn is not None: - del self.attn - self.attn = None + if hasattr(self, 'pre_dispatch_computation') and self.pre_dispatch_computation is not None: + del self.pre_dispatch_computation + self.pre_dispatch_computation = None if hasattr(self, 'moe_dispatch') and self.moe_dispatch is not None: del self.moe_dispatch self.moe_dispatch = None @@ -114,23 +117,20 @@ def release_state(self): def _build_callable_nodes(self, event, comp_stream, comm_stream, extra_args): """ Builds the callable nodes for the transformer/mtp layer: - attn, mlp, moe_dispatch, moe_combine, and mtp_post_process. + pre_dispatch_computation, moe_dispatch, mlp, moe_combine, + and mtp_post_process. """ - from megatron.core.models.gpt.fine_grained_callables import ( - TransformerLayerNode, + from megatron.core.models.common.fine_grained_callables import ( build_layer_callables, + get_layer_moe_metadata, ) - from megatron.core.transformer.moe.moe_layer import MoELayer + from megatron.core.models.common.utils import TransformerLayerNode from megatron.core.transformer.multi_token_prediction import MultiTokenPredictionLayer - # build the forward and backward callables for the transformer/mtp layer fwd_callables, bwd_dw_callable_map = build_layer_callables(self.layer) + is_moe, num_local_experts = get_layer_moe_metadata(self.layer) - # get flags for latter use is_mtp = isinstance(self.layer, MultiTokenPredictionLayer) - transformer_layer = self.layer.mtp_model_layer if is_mtp else self.layer - is_moe = isinstance(transformer_layer.mlp, MoELayer) - num_local_experts = transformer_layer.mlp.num_local_experts if is_moe else None extra_args["config"] = self.layer.config extra_args["is_moe"] = is_moe @@ -153,7 +153,7 @@ def create_node(stream, module, name): ) ( - attn_module, + pre_dispatch_module, moe_dispatch_module, mlp_module, moe_combine_module, @@ -162,7 +162,9 @@ def create_node(stream, module, name): # Create nodes for different operations in the layer # Each node type has a predefined name that determines its memory strategy - self.attn = create_node(comp_stream, attn_module, "attn") + self.pre_dispatch_computation = create_node( + comp_stream, pre_dispatch_module, "pre_dispatch_computation" + ) self.mlp = create_node(comp_stream, mlp_module, "mlp") if is_moe: self.moe_dispatch = create_node(comm_stream, moe_dispatch_module, "moe_dispatch") @@ -211,21 +213,25 @@ def set_fsdp_reshard_hooks(self, post_forward_hook, post_backward_hook): post_backward_hook: Callable(module) that releases backward-pass params (bwd=True). Typically ``fsdp_wrapper.post_backward_release_module``. """ + from megatron.core.models.hybrid.hybrid_block import HybridStack from megatron.core.transformer.multi_token_prediction import MultiTokenPredictionLayer from megatron.core.transformer.transformer_layer import TransformerLayer - assert isinstance(self.layer, (TransformerLayer, MultiTokenPredictionLayer)), ( + assert isinstance(self.layer, (TransformerLayer, HybridStack, MultiTokenPredictionLayer)), ( f"Megatron FSDP with EP Overlap only supports TransformerLayer, " + f"HybridStack and MultiTokenPredictionLayer, " f"but got {type(self.layer).__name__}." ) - if isinstance(self.layer, TransformerLayer): + if isinstance(self.layer, (TransformerLayer, HybridStack)): hook_module = self.layer else: hook_module = self.layer.mtp_model_layer - # After the last backward op (attn), release backward-pass params. - self.attn.set_post_backward_hook(lambda: post_backward_hook(hook_module)) + # After the last backward op (pre_dispatch_computation), release backward-pass params. + self.pre_dispatch_computation.set_post_backward_hook( + lambda: post_backward_hook(hook_module) + ) # Determine the last node in forward order. if isinstance(self.moe_combine, NoopScheduleNode): @@ -254,12 +260,12 @@ def run(f_layer, b_layer, f_input=None, b_grad=None, is_last_layer_in_bwd=False) """Schedule one-forward-one-backward operations for a single transformer layer. This function interleaves forward and backward operations, overlapping the communications - (dispatch or combine) of one with the computations (att or mlp) of the other + (dispatch or combine) of one with the computations (pre_dispatch or mlp) of the other to maximize parallelism and efficiency. When f_layer and b_layer are not None, forward and backward pass are overlapped as follows: - comm_stream: combine_bwd | dispatch_fwd->dispatch_bwd | combine_fwd - comp_stream: attn_fwd | mlp_bwd->mlp_bwd_dw->mlp_fwd| attn_bwd + comm_stream: combine_bwd | dispatch_fwd->dispatch_bwd | combine_fwd + comp_stream: pre_dispatch_fwd | mlp_bwd->mlp_bwd_dw->mlp_fwd| pre_dispatch_bwd MLP-side mHC post-processing runs inside the combine node on the communication stream. Group recompute runs on the normal compute stream immediately before the node containing mHC post-processing backward. @@ -286,7 +292,7 @@ def run(f_layer, b_layer, f_input=None, b_grad=None, is_last_layer_in_bwd=False) if f_layer is not None: with f_layer.get_fp8_context(): - f_input = f_layer.attn.forward(f_input) + f_input = f_layer.pre_dispatch_computation.forward(f_input) if b_layer is not None: b_grad = b_layer.mlp.backward(b_grad) @@ -300,7 +306,7 @@ def run(f_layer, b_layer, f_input=None, b_grad=None, is_last_layer_in_bwd=False) b_grad = b_layer.moe_dispatch.backward(b_grad) if b_layer is not None and b_layer.config.ep_overlap_early_attn_memory_release: - b_grad = b_layer.attn.backward(b_grad) + b_grad = b_layer.pre_dispatch_computation.backward(b_grad) if f_layer is not None: with f_layer.get_fp8_context(): @@ -311,16 +317,16 @@ def run(f_layer, b_layer, f_input=None, b_grad=None, is_last_layer_in_bwd=False) f_input = f_layer.moe_combine.forward(f_input) if b_layer is not None and not b_layer.config.ep_overlap_early_attn_memory_release: - b_grad = b_layer.attn.backward(b_grad) + b_grad = b_layer.pre_dispatch_computation.backward(b_grad) if f_layer is not None: with f_layer.get_fp8_context(): f_input = f_layer.mtp_post_process.forward(f_input) - # Delay the last attn_dw in backward pass (attn_dw of the first layer) - # for overlapping with the p2p comm + # Delay the last pre_dispatch_computation wgrad in backward pass (wgrad + # of the first layer) for overlapping with the p2p comm. if b_layer is not None and not is_last_layer_in_bwd: - b_layer.attn.backward_dw() + b_layer.pre_dispatch_computation.backward_dw() return f_input, b_grad @@ -338,8 +344,27 @@ class TransformerModelChunkSchedulePlan(AbstractSchedulePlan): │ ├── layer[1]: TransformerLayerSchedulePlan │ └── ... └── post_process: PostProcessNode + + Subclasses can swap the per-layer schedule plan by overriding the + ``LAYER_SCHEDULE_PLAN_CLASS`` class attribute (e.g. HybridStack uses a + layer plan that understands grouped/inferred layer types). They can also + swap the pre/post-process node classes via ``PRE_PROCESS_NODE_CLASS`` / + ``POST_PROCESS_NODE_CLASS`` so each model owns its own embedding / output + layer node implementations. """ + #: The TransformerLayerSchedulePlan-compatible class used to build per-layer + #: schedule plans. Subclasses override this to inject a layer-plan variant. + LAYER_SCHEDULE_PLAN_CLASS = None + + #: Pre/post-process node classes. Defaults below pull in the GPT-side + #: ``PreProcessNode`` / ``PostProcessNode`` (which call ``GPTModel._preprocess`` / + #: ``GPTModel._postprocess``). Subclasses set these to model-specific node + #: classes so the node calls the right model's ``_preprocess`` / + #: ``_postprocess`` methods. + PRE_PROCESS_NODE_CLASS = None + POST_PROCESS_NODE_CLASS = None + def __init__( self, model, @@ -381,7 +406,10 @@ def __init__( Returns: The model chunk schedule plan. """ - from megatron.core.models.gpt.fine_grained_callables import PostProcessNode, PreProcessNode + from megatron.core.models.common.utils import PostProcessNode, PreProcessNode + + pre_process_cls = self.PRE_PROCESS_NODE_CLASS or PreProcessNode + post_process_cls = self.POST_PROCESS_NODE_CLASS or PostProcessNode self._model_chunk_state = ModelChunkState() self._transformer_layers = [] @@ -411,7 +439,7 @@ def __init__( self._model_chunk_state.attention_bias = None # build preprocess - self.pre_process = PreProcessNode( + self.pre_process = pre_process_cls( model, self._model_chunk_state, self._event, get_comp_stream ) @@ -427,7 +455,7 @@ def __init__( # build post process if model.post_process: - self.post_process = PostProcessNode( + self.post_process = post_process_cls( model, self._model_chunk_state, self._event, get_comp_stream ) @@ -437,6 +465,7 @@ def _build_layer_schedule_plan(self, module, comp_stream, comm_stream, module_ta from megatron.core.tensor_parallel.random import CheckpointManager + plan_cls = self.LAYER_SCHEDULE_PLAN_CLASS or TransformerLayerSchedulePlan num_layers = len(module.layers) config = module.config use_mhc_recompute = ( @@ -457,15 +486,16 @@ def _build_layer_schedule_plan(self, module, comp_stream, comm_stream, module_ta mhc_recompute_manager is not None and (layer_idx == num_layers - 1 or (layer_idx + 1) % group_size == 0) ) - extra_args = { - "is_first_layer": layer_idx == 0, - "is_last_layer": layer_idx == num_layers - 1, - "mhc_recompute_manager": mhc_recompute_manager, - "is_last_layer_in_mhc_recompute_group": is_group_end, - "mhc_recompute_group_index": group_index, - "mhc_recompute_module_tag": module_tag, - } - layer_plan = TransformerLayerSchedulePlan( + extra_args = self._extra_args_for_layer(module, layer_idx, num_layers) + extra_args.update( + { + "mhc_recompute_manager": mhc_recompute_manager, + "is_last_layer_in_mhc_recompute_group": is_group_end, + "mhc_recompute_group_index": group_index, + "mhc_recompute_module_tag": module_tag, + } + ) + layer_plan = plan_cls( module.layers[layer_idx], self.event, self.state, @@ -479,6 +509,14 @@ def _build_layer_schedule_plan(self, module, comp_stream, comm_stream, module_ta group_index += 1 mhc_recompute_manager = CheckpointManager() + def _extra_args_for_layer(self, module, layer_idx, num_layers): + """Per-layer ``extra_args`` dict passed to the layer plan constructor. + + Subclasses extend this hook to thread additional metadata (e.g. hybrid + layer-type symbols) without overriding ``_build_layer_schedule_plan``. + """ + return {"is_first_layer": layer_idx == 0, "is_last_layer": layer_idx == num_layers - 1} + @property def event(self): """Gets the CUDA event for synchronization.""" @@ -621,22 +659,22 @@ def run( if f_schedule_plan is not None and post_forward is not None: # post_forward()/send_forward_recv_forward() is running in the communication stream, - # so the p2p comm could be overlapped with the attn backward + # so the p2p comm could be overlapped with the pre_dispatch backward with torch.cuda.stream(get_comm_stream()): f_schedule_plan.wait_current_stream() post_forward(f_input, f_schedule_plan.vp_stage) # post_backward()/send_backward_recv_backward() is running in the computation stream, - # so the p2p comm could be overlapped with the wgrad of attn backward + # so the p2p comm could be overlapped with the wgrad of pre_dispatch backward if b_schedule_plan is not None and post_backward is not None: b_schedule_plan.wait_current_stream() post_backward(b_grad, b_schedule_plan.vp_stage) - # Delay the last attn_dw in backward pass (attn_dw of the first layer) - # for overlapping with the p2p comm + # Delay the last pre_dispatch_computation wgrad in backward pass (wgrad + # of the first layer) for overlapping with the p2p comm. if b_num_layers > 0: assert b_layer is not None - b_layer.attn.backward_dw() + b_layer.pre_dispatch_computation.backward_dw() b_layer.release_state() # post process forward diff --git a/megatron/core/models/common/utils.py b/megatron/core/models/common/utils.py new file mode 100644 index 00000000000..d1b7c54fefb --- /dev/null +++ b/megatron/core/models/common/utils.py @@ -0,0 +1,437 @@ +# Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Schedule-plan helpers shared by GPTModel and HybridModel. + +These pieces used to live in ``core/models/gpt/fine_grained_callables.py`` and +were imported by ``core/models/common/model_chunk_schedule_plan.py`` and the +hybrid schedule plan via that path. They are model-agnostic in practice — the +``Pre/PostProcessNode`` classes call the model's ``_preprocess`` / +``_postprocess`` methods and don't otherwise care which model implements +them — so they live here now. +""" + +import weakref +from functools import partial +from typing import Callable + +import torch + +from megatron.core.pipeline_parallel.utils import ScheduleNode, make_viewless +from megatron.core.transformer.enums import CudaGraphModule +from megatron.core.transformer.module import GraphableMegatronModule, float16_to_fp32 +from megatron.core.transformer.transformer_layer import TransformerLayer, make_viewless_tensor +from megatron.core.utils import internal_api, nvtx_range_pop, nvtx_range_push + + +def weak_method(method): + """Wrap ``method`` in a weakref-keyed dispatcher to break refcycles. + + ``ScheduleNode`` keeps a reference to the bound forward / backward functions + of every node in the plan; using a strong reference would keep the layer + plan (and the model chunk through it) alive after the iteration completes. + The ``weakref.WeakMethod`` lets the schedule plan be torn down between + iterations without manual ``del`` chains. + """ + method_ref = weakref.WeakMethod(method) + del method + + def wrapped_func(*args, **kwarg): + return method_ref()(*args, **kwarg) + + return wrapped_func + + +@internal_api +def should_free_input(name, is_moe, config, num_local_experts): + """Whether the schedule node named ``name`` can free its input after forward. + + The schedule decomposes a transformer layer into ``pre_dispatch_computation``, + ``moe_dispatch``, ``mlp``, and ``moe_combine`` nodes; the inputs to some of + those nodes are not needed in backward and can be released early to lower + peak activation memory. Dense layers and the ``pre_dispatch_computation`` + node always need their input retained (the attention residual flows through + the post-MLP BDA). + + Args: + name: Schedule node name. + is_moe: True for MoE layers; dense layers always retain inputs. + config: ``TransformerConfig`` for the layer. + num_local_experts: Local expert count on this rank (None for dense). + + Returns: + True iff the named node may free its input after forward. + """ + # For dense layers [pre_dispatch_computation, fake, mlp, fake], the input is needed + # during backward pass + if not is_moe: + return False + enable_deepep = ( + config.moe_token_dispatcher_type == "flex" + and config.moe_flex_dispatcher_backend == "deepep" + ) + enable_hybridep = ( + config.moe_token_dispatcher_type == "flex" + and config.moe_flex_dispatcher_backend == "hybridep" + ) + enable_ncclep = ( + config.moe_token_dispatcher_type == "flex" + and config.moe_flex_dispatcher_backend == "ncclep" + ) + # Define which nodes should free input memory. + # Since we split the computing graph into multiple nodes, we can manually control + # when and how to free the input memory. + # The input and output of A2A are not needed anymore after the forward pass, + # so we can free the input memory after the forward pass. + + # When low precision fp8/4 is enabled, the casted tensors are saved and the + # original bf16 tensors are safe to be freed. + free_mlp = config.fp8 is not None or config.fp4 is not None + if not free_mlp: + # AlltoAll dispatcher with local_num_experts=1, HybridEP, and NCCL EP all use + # identity operation for `dispatch_postprocess`, hence the mlp inputs will be + # directly passed to GroupedGemm and should be saved for backward pass. + free_mlp = num_local_experts > 1 or config.moe_token_dispatcher_type != "alltoall" + free_mlp = free_mlp and not (enable_hybridep or enable_ncclep) + + free_input_nodes = { + "mlp": free_mlp, + "moe_combine": True, + # For non-DeepEP/HybridEP/NCCL-EP dispatcher mode, the input is the un-dispatched + # tokens and probs before dispatch A2A and it's not needed anymore after the + # forward pass. For DeepEP, HybridEP, and NCCL EP dispatcher mode, they are all + # needed in backward pass and cannot be freed. + # If moe_preprocess is in cuda graph scope, tokens and probs are fixed size + # tensors, so they cannot be freed. + "moe_dispatch": not (enable_deepep or enable_hybridep or enable_ncclep) + and (CudaGraphModule.moe_preprocess not in config.cuda_graph_modules), + } + + return free_input_nodes.get(name, False) + + +class LayerState: + """State shared between the schedule nodes that come from one logical layer. + + Empty placeholder; nodes attach their own attributes (residual, dispatched + probs, shared-expert outputs) for downstream nodes in the same layer to + consume. Kept as a real class so weakrefs work uniformly. + """ + + pass + + +class PreProcessNode(ScheduleNode): + """Run the model's ``_preprocess`` (embedding + rotary + padding mask). + + The schedule plan wraps a model that exposes a ``_preprocess`` method + returning the canonical 6-tuple ``(decoder_input, rotary_pos_emb, + rotary_pos_cos, rotary_pos_sin, sequence_len_offset, padding_mask)`` + (slots a given model doesn't use are returned as ``None``). The chunk + state is mutated in-place so layer nodes can read the same fields by + name. + """ + + def __init__(self, model, chunk_state, event, stream): + super().__init__(weak_method(self.forward_impl), stream, event, name="pre_process") + self.model = model + self.chunk_state = chunk_state + + def forward_impl(self): + """Run model preprocessing and store chunk-level inputs for layer nodes.""" + if not self.model.pre_process: + self.chunk_state.decoder_input = self.model.decoder.input_tensor + ( + decoder_input, + rotary_pos_emb, + rotary_pos_cos, + rotary_pos_sin, + sequence_len_offset, + padding_mask, + ) = self.model._preprocess( + input_ids=self.chunk_state.input_ids, + position_ids=self.chunk_state.position_ids, + decoder_input=self.chunk_state.decoder_input, + packed_seq_params=self.chunk_state.packed_seq_params, + padding_mask=self.chunk_state.padding_mask, + ) + + self.chunk_state.decoder_input = decoder_input + self.chunk_state.rotary_pos_emb = rotary_pos_emb + self.chunk_state.rotary_pos_cos = rotary_pos_cos + self.chunk_state.rotary_pos_sin = rotary_pos_sin + self.chunk_state.sequence_len_offset = sequence_len_offset + self.chunk_state.padding_mask = padding_mask + return decoder_input + + +class PostProcessNode(ScheduleNode): + """Run the model's ``_postprocess`` (final norm, output layer, loss). + + Calls ``_postprocess`` with ``mtp_in_postprocess=False`` because the + schedule plan handles MTP layers as sibling layer nodes inside the same + chunk; the model's MTP block is not invoked here. The optional final + layernorm — applied only when this rank holds an empty decoder shard + (early stage of pipeline parallel) — is handled here so the chunk plan + does not need a separate node for it. + """ + + def __init__(self, model, chunk_state, event, stream): + super().__init__(weak_method(self.forward_impl), stream, event, name="post_process") + self.model = model + self.chunk_state = chunk_state + + def forward_impl(self, hidden_states): + """Run model postprocessing for the chunk's final hidden states.""" + empty_decoder = len(self.model.decoder.layers) == 0 + layer_norm = getattr(self.model.decoder, "final_norm", None) or getattr( + self.model.decoder, "final_layernorm", None + ) + if not self.model.config.mtp_num_layers and empty_decoder and layer_norm: + hidden_states = layer_norm(hidden_states) + hidden_states = make_viewless_tensor( + inp=hidden_states, requires_grad=True, keep_graph=True + ) + + loss = self.model._postprocess( + hidden_states=hidden_states, + input_ids=self.chunk_state.input_ids, + position_ids=self.chunk_state.position_ids, + labels=self.chunk_state.labels, + decoder_input=self.chunk_state.decoder_input, + rotary_pos_emb=self.chunk_state.rotary_pos_emb, + rotary_pos_cos=self.chunk_state.rotary_pos_cos, + rotary_pos_sin=self.chunk_state.rotary_pos_sin, + mtp_in_postprocess=False, + loss_mask=self.chunk_state.loss_mask, + attention_mask=self.chunk_state.attention_mask, + packed_seq_params=self.chunk_state.packed_seq_params, + sequence_len_offset=self.chunk_state.sequence_len_offset, + runtime_gather_output=self.chunk_state.runtime_gather_output, + extra_block_kwargs=self.chunk_state.extra_block_kwargs, + output_processor=self.chunk_state.output_processor, + output_processor_context=self.chunk_state.output_processor_context, + ) + + # combined-1F1B currently expects fp32 loss output. + return float16_to_fp32(loss) + + +class TransformerLayerNode(ScheduleNode): + """Schedule node for one slot of a fine-grained transformer layer plan. + + Each transformer layer is decomposed into ``pre_dispatch_computation``, + ``moe_dispatch``, ``mlp``, and ``moe_combine`` slots; this class is the scheduler-side + handle for one slot. It owns the slot's stream / event, the per-slot + ``free_input`` policy, and the optional delayed weight-gradient hook. + Subclasses override ``_resolve_free_input`` to specialize the policy + (HybridStackNode does this for grouped layers). + """ + + def __init__( + self, + stream, + event, + layer_state, + chunk_state, + submodule, + name="default", + bwd_dw_callables=None, + extra_args={}, + ): + config = extra_args.get("config", None) + assert config is not None, "model config must be passed to TransformerLayerNode." + is_moe = extra_args.get("is_moe", False) + num_local_experts = extra_args.get("num_local_experts", None) + free_input = self._resolve_free_input(name, is_moe, config, num_local_experts) + self.delay_wgrad_compute = extra_args.get("delay_wgrad_compute", False) + + super().__init__( + weak_method(self.forward_impl), + stream, + event, + weak_method(self.backward_impl), + free_input=free_input, + name=name, + ncclep_zero_copy=config.moe_ncclep_zero_copy, + ) + self.layer_state = layer_state + self.chunk_state = chunk_state + self.submodule = submodule + self.detached = tuple() + self.before_detached = tuple() + self.is_mtp = extra_args.get("is_mtp", False) + self.post_wgrad_grad_acc_hooks = None + + self.is_first_layer = extra_args.get("is_first_layer", False) + self.is_last_layer = extra_args.get("is_last_layer", False) + + # Whether this slot is the first/last node of its TransformerLayer in + # forward / backward order. Set by ``set_post_*_hook``; used to decide + # when to invoke the layer-level FSDP reshard hooks. + self.is_layer_first_node = None + self.is_layer_last_node = None + + self.bwd_dw_callables = [] + if bwd_dw_callables is not None: + self.bwd_dw_callables = ( + bwd_dw_callables if isinstance(bwd_dw_callables, list) else [bwd_dw_callables] + ) + + @staticmethod + def _resolve_free_input(name, is_moe, config, num_local_experts): + """Free-input policy hook. Subclasses override to specialize.""" + return should_free_input(name, is_moe, config, num_local_experts) + + def detach(self, t): + """Detach a tensor and remember it for backward through the schedule node.""" + detached = make_viewless(t).detach() + detached.requires_grad = t.requires_grad + self.before_detached = self.before_detached + (t,) + self.detached = self.detached + (detached,) + return detached + + def forward_impl(self, *args): + """Invoke the slot's submodule forward.""" + return self.submodule(self, *args) + + def backward_impl(self, outputs, output_grad): + """Run the slot's backward and return the input grads.""" + detached_grad = tuple([e.grad for e in self.detached]) + grads = output_grad + detached_grad + self.default_backward_func(outputs + self.before_detached, grads) + + return grads + + def forward(self, *inputs): + """Execute forward and fire the per-layer post-forward hook on the last slot.""" + output = super().forward(*inputs) + if self.is_layer_last_node: + self._post_forward_hook() + return output + + def backward(self, *output_grad): + """Execute backward and fire the per-layer post-backward hook on the first slot. + + When ``delay_wgrad_compute`` is set, the hook fires after ``backward_dw`` + instead, because the wgrad work has not yet run when ``backward`` returns. + """ + grads = super().backward(*output_grad) + if not self.delay_wgrad_compute and self.is_layer_first_node: + self._post_backward_hook() + return grads + + def backward_dw(self): + """Run the slot's delayed weight-gradient callables on the slot's stream.""" + if not self.delay_wgrad_compute: + return + if isinstance(self.stream, Callable): + self.stream = self.stream() + with torch.cuda.stream(self.stream): + nvtx_msg = f"{self.name} wgrad" + nvtx_range_push(nvtx_msg) + for module in self.bwd_dw_callables: + module.backward_dw() + nvtx_range_pop(nvtx_msg) + + # Collect ``post_wgrad_grad_acc_hook`` from params whose grads were + # produced by *this* slot's wgrad callables. The hook must run on the + # same stream right after the wgrad it depends on; collecting on the + # first invocation makes the per-iteration hook order deterministic. + if self.post_wgrad_grad_acc_hooks is None: + self.post_wgrad_grad_acc_hooks = [] + for module in self.bwd_dw_callables: + for param in module.parameters(): + if ( + getattr(param, "post_wgrad_grad_acc_hook", False) + and param.requires_grad + and param.grad is not None + ): + self.post_wgrad_grad_acc_hooks.append(param.post_wgrad_grad_acc_hook) + + if self.post_wgrad_grad_acc_hooks: + with torch.cuda.stream(self.stream): + for hook in self.post_wgrad_grad_acc_hooks: + hook() + + if self.is_layer_first_node: + self._post_backward_hook() + self.bwd_dw_callables = None + + def set_post_forward_hook(self, hook): + """Mark this slot as the layer's last fwd node and register the hook.""" + self.is_layer_last_node = True + self._post_forward_hook = hook + + def set_post_backward_hook(self, hook): + """Mark this slot as the layer's first bwd node and register the hook.""" + self.is_layer_first_node = True + self._post_backward_hook = hook + + def __del__(self): + # Release references early to help avoid leaks across iterations. + self.before_detached = None + self.detached = None + self.layer_state = None + self.chunk_state = None + self.submodule = None + + +class _BackwardDWWrapper: + """Backward weight-gradient wrapper for a transformer pre-dispatch slot. + + Runs the layer's ``self_attention.backward_dw`` plus, on MoE layers, the + shared-expert ``backward_dw``; coordinates with the cuda-graph wgrad + capture (``set_graphed_backward_dw_callable``) so that scopes covered by + the graph are not re-run eagerly. Used when + ``overlap_moe_expert_parallel_comm`` and ``delay_wgrad_compute`` are both + enabled. + """ + + def __init__(self, layer): + assert isinstance( + layer, GraphableMegatronModule + ), "cuda graphed ep overlap only supports GraphableMegatronModule." + assert isinstance( + layer, TransformerLayer + ), "cuda graphed ep overlap only supports TransformerLayer for now." + self.layer = layer + self.graphed_backward_dw_callable = None + self.attn_dw_callable = layer.self_attention.backward_dw + self.submodules = [layer.self_attention] + if layer.is_moe_layer: + self.shared_expert_dw_callable = partial( + layer.mlp.backward_dw, routed_experts=False, shared_experts=True + ) + if layer.mlp.use_shared_expert: + self.submodules.append(layer.mlp.shared_experts) + else: + self.shared_expert_dw_callable = None + self.cuda_graph_modules = layer.config.cuda_graph_modules + + def backward_dw(self): + """Run eager or graphed backward wgrad callables for the wrapped layer.""" + is_replay = hasattr(self.layer, 'cuda_graphs') and self.layer.cuda_graphs + if self.shared_expert_dw_callable is not None and ( + not is_replay or CudaGraphModule.moe_router not in self.cuda_graph_modules + ): + self.shared_expert_dw_callable() + if not is_replay or CudaGraphModule.attn not in self.cuda_graph_modules: + self.attn_dw_callable() + if is_replay and self.graphed_backward_dw_callable is not None: + self.graphed_backward_dw_callable() + self.layer = None + + def set_graphed_backward_dw_callable(self, graphed_backward_dw_callable): + """Plug the cuda-graph backward wgrad replay callable.""" + self.graphed_backward_dw_callable = graphed_backward_dw_callable + + def parameters(self): + """Yield parameters from the wrapped layer's wgrad submodules. + + Mirrors ``torch.nn.Module.parameters`` so callers (notably + ``TransformerLayerNode.backward_dw``) can collect ``post_wgrad_grad_acc_hook`` + without knowing the concrete layer layout. + """ + for module in self.submodules: + for param in module.parameters(): + yield param diff --git a/megatron/core/models/gpt/fine_grained_callables.py b/megatron/core/models/gpt/fine_grained_callables.py index 34533ada4e6..3aaa0dac32a 100644 --- a/megatron/core/models/gpt/fine_grained_callables.py +++ b/megatron/core/models/gpt/fine_grained_callables.py @@ -13,7 +13,7 @@ from megatron.core.pipeline_parallel.fine_grained_activation_offload import ( FineGrainedActivationOffloadingInterface as off_interface, ) -from megatron.core.pipeline_parallel.utils import ScheduleNode, make_viewless +from megatron.core.pipeline_parallel.utils import ScheduleNode, StageDispatchBwdGrad, make_viewless from megatron.core.transformer.enums import CudaGraphModule from megatron.core.transformer.module import GraphableMegatronModule, float16_to_fp32 from megatron.core.transformer.moe.moe_layer import MoELayer @@ -530,12 +530,16 @@ def build_transformer_layer_callables(layer: TransformerLayer): functions. This decomposition separates computation-heavy tasks (e.g., self-attention, MLP) from communication-heavy tasks (e.g., MoE's All-to-All). - The five callable slots are: - 1. Attention and routing preprocess (computation) - 2. MoE Dispatch (communication) - 3. MLP / MoE Experts (computation) - 4. MoE Combine and MLP-side mHC post-processing (communication) - 5. MTP post-processing (computation, MTP layers only) + The five callables align with the schedule plan's slot order: + 1. pre_dispatch_computation (computation): + attention -> pre-MLP layernorm -> router -> dispatch preprocess. + For dense layers this is just the attention pass. + 2. moe_dispatch (communication): MoE dispatch All-to-All. + 3. mlp / moe_experts (computation): dense MLP or routed-experts compute. + 4. moe_combine (communication): MoE combine All-to-All + post-MLP residual, + including MLP-side mHC post-processing. + 5. mtp_post_process (computation): always ``None`` here; only the MTP + wrapper in ``common/fine_grained_callables.py`` fills this slot. By assigning these functions to different CUDA streams (e.g., a compute stream and a communication stream), the scheduler can overlap their execution, preventing @@ -547,8 +551,11 @@ def build_transformer_layer_callables(layer: TransformerLayer): Returns: A tuple containing: - - forward_funcs: List of callable functions for the layer - - backward_dw: Dict of weight gradient functions for the layer + - forward_funcs: List of 5 callables, one per slot in the schedule plan + (pre_dispatch_computation, moe_dispatch, mlp, moe_combine, + mtp_post_process=None). + - backward_dw: Dict mapping slot name to the delayed-wgrad callable + (keys: "pre_dispatch_computation", "mlp"). """ is_moe = isinstance(layer.mlp, MoELayer) enable_deepep = ( @@ -566,9 +573,9 @@ def build_transformer_layer_callables(layer: TransformerLayer): is_hyper_connection_layer = isinstance(layer, HyperConnectionTransformerLayer) is_mhc_layer = is_moe and is_hyper_connection_layer - def submodule_attn_forward(node: ScheduleNode, hidden_states: torch.Tensor): + def submodule_pre_dispatch_forward(node: ScheduleNode, hidden_states: torch.Tensor): """ - Performs same attnention forward logic as GPT Model and forward pass for + Performs the same attention forward logic as GPTModel and the forward pass for computations between attention and dispatch: pre mlp layernorm->router->dispatch preprocess """ @@ -719,11 +726,17 @@ def submodule_dispatch_forward( token_dispatcher = layer.mlp.token_dispatcher if enable_deepep or enable_hybridep or enable_ncclep: # update token_probs to be the detached version, prevents - # backward graph from connecting to attn submodule + # backward graph from connecting to pre_dispatch_computation submodule token_dispatcher._comm_manager.token_probs = probs dispatched_tokens, dispatched_probs = layer.mlp.dispatch(local_tokens, probs) + if enable_ncclep and layer.config.moe_ncclep_zero_copy: + # Insert an identity node as the sole consumer of the dispatch output, so the + # dispatch-backward gets the symm buffer instead of a non-symm AccumulateGrad clone. + # Must stay inside this node's graph segment (before the next node detaches it). + dispatched_tokens = StageDispatchBwdGrad.apply(dispatched_tokens, token_dispatcher) + # `dispatched_probs` is needed by backward pass of swiglu, therefore it's # passed to moe_forward within `layer_state` to avoid the free_input process # of the input tensors. @@ -861,15 +874,15 @@ def raise_not_implemented(*args): raise NotImplementedError("This callable is not implemented for Dense layer.") # Build forward and backward callable functions - attn_func = submodule_attn_forward + pre_dispatch_func = submodule_pre_dispatch_forward dispatch_func = submodule_dispatch_forward if is_moe else raise_not_implemented mlp_func = submodule_moe_forward if is_moe else mlp_wrapper combine_func = submodule_combine_forward if is_moe else raise_not_implemented layer.init_backward_dw_wrapper() - forward_funcs = [attn_func, dispatch_func, mlp_func, combine_func, None] - backward_dw = {"attn": layer.backward_dw_wrapper, "mlp": layer.mlp} + forward_funcs = [pre_dispatch_func, dispatch_func, mlp_func, combine_func, None] + backward_dw = {"pre_dispatch_computation": layer.backward_dw_wrapper, "mlp": layer.mlp} return forward_funcs, backward_dw diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index a986463bcde..41aa6bd88cd 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -1,5 +1,6 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +import logging from collections import OrderedDict from typing import Any, Callable, Dict, Literal, Optional @@ -42,8 +43,11 @@ WrappedTensor, deprecate_inference_params, is_using_quantization_scales, + log_single_rank, ) +logger = logging.getLogger(__name__) + class GPTModel(LanguageModule): """GPT Transformer language model. @@ -112,6 +116,13 @@ def __init__( pg_collection: Optional[ProcessGroupCollection] = None, vp_stage: Optional[int] = None, ) -> None: + log_single_rank( + logger, + logging.WARNING, + "GPTModel IS DEPRECATED. GPTModel is only accepting critical bug fixes, no new " + "features. Please reference the migration guide " + "`docs/user-guide/hybrid-model-migration.md` for details on how to use `HybridModel`", + ) super().__init__(config=config, pg_collection=pg_collection) if has_config_logger_enabled(config): diff --git a/megatron/core/models/hybrid/hybrid_block.py b/megatron/core/models/hybrid/hybrid_block.py index 297bfc4b054..e2ba344bc22 100644 --- a/megatron/core/models/hybrid/hybrid_block.py +++ b/megatron/core/models/hybrid/hybrid_block.py @@ -5,6 +5,7 @@ # This source code is licensed under the Apache license found in the # LICENSE file in the root directory of this source tree. +import copy from contextlib import nullcontext from dataclasses import dataclass from typing import List, Optional, Tuple, Union @@ -15,7 +16,7 @@ from megatron.core.dist_checkpointing.mapping import ShardedStateDict from megatron.core.dist_checkpointing.utils import replace_prefix_for_sharding from megatron.core.enums import Fp8Recipe -from megatron.core.extensions.transformer_engine import TENorm +from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear, TENorm from megatron.core.fp4_utils import get_fp4_context from megatron.core.fp8_utils import get_fp8_context from megatron.core.inference.contexts import BaseInferenceContext @@ -39,6 +40,7 @@ convert_module_to_dtype_except_fp32_marked, mark_keep_in_fp32, ) +from megatron.core.transformer.multi_latent_attention import FusedMLASelfAttention from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.transformer_layer import TransformerLayer from megatron.core.transformer.utils import ( @@ -62,6 +64,7 @@ class HybridStackSubmodules: csa_layer: Union[ModuleSpec, type] = IdentityOp hca_layer: Union[ModuleSpec, type] = IdentityOp window_layer: Union[ModuleSpec, type] = IdentityOp + mla_layer: Union[ModuleSpec, type] = IdentityOp mlp_layer: Union[ModuleSpec, type] = IdentityOp moe_layer: Union[ModuleSpec, type] = IdentityOp mtp_block_spec: Optional[ModuleSpec] = None @@ -594,6 +597,9 @@ def __init__( ) self.layer_type_list = layer_type_list + if getattr(self.config, "mla_down_proj_fusion", False): + submodules = self._fuse_mla_down_proj(submodules) + # Build layers from the pre-selected segment self.layers = nn.ModuleList() for i, layer_type in enumerate(self.layer_type_list): @@ -670,6 +676,16 @@ def __init__( add_layer_offset=False, pp_layer_offset=pp_layer_offset, ) + elif layer_type == LayerSymbols.MLA: + layer = build_module( + submodules.mla_layer, + config=self.config, + layer_number=layer_number, + pg_collection=pg_collection, + is_mtp_layer=is_mtp_layer, + add_layer_offset=False, + pp_layer_offset=pp_layer_offset, + ) elif layer_type == LayerSymbols.MLP: layer = build_module( submodules.mlp_layer, @@ -747,6 +763,26 @@ def _set_mtp_layer_number_for_moe_metrics( if router is not None and getattr(router, "is_mtp_layer", False): router.mtp_layer_number = mtp_layer_number + def _fuse_mla_down_proj(self, submodules: HybridStackSubmodules) -> HybridStackSubmodules: + # Avoid modifying the original object so users don't get surprised about their `submodules` + # being modified underneath them. + submodules = copy.deepcopy(submodules) + mla_spec = submodules.mla_layer + # We always fuse the input layernorm because Hybrid always uses TransformerEngine. + mla_spec.submodules.input_layernorm = IdentityOp + mla_spec.submodules.self_attention.module = FusedMLASelfAttention + mla_spec.submodules.self_attention.submodules.linear_qkv_down_proj = ( + TELayerNormColumnParallelLinear + ) + mla_spec.submodules.self_attention.submodules.linear_q_down_proj = None + mla_spec.submodules.self_attention.submodules.linear_kv_down_proj = None + mla_spec.submodules.sharded_state_dict_keys_map = { + "self_attention.linear_q_down_proj.layer_norm_": "input_layernorm.", + "self_attention.linear_kv_down_proj.layer_norm_": "input_layernorm.", + "self_attention.linear_qkv_down_proj.layer_norm_": "input_layernorm.", + } + return submodules + def set_input_tensor(self, input_tensor: Tensor): """Set input tensor to be used instead of forward()'s input. diff --git a/megatron/core/models/hybrid/hybrid_layer_allocation.py b/megatron/core/models/hybrid/hybrid_layer_allocation.py index a8d2006c3b0..90e12b40d14 100644 --- a/megatron/core/models/hybrid/hybrid_layer_allocation.py +++ b/megatron/core/models/hybrid/hybrid_layer_allocation.py @@ -21,13 +21,14 @@ class Symbols: CSA = "C" # DSv4 Compressed Sparse Attention (compress_ratio=4) HCA = "H" # DSv4 Heavily Compressed Attention (compress_ratio=128) WINDOW = "W" # DSv4 sliding-window-only attention (compress_ratio=0; no compressor/indexer) + MLA = "+" MLP = "-" MOE = 'E' PIPE = '|' MTP_SEPARATOR = "/" - VALID_LAYERS = {MAMBA, GDN, ATTENTION, DS_ATTENTION, CSA, HCA, WINDOW, MLP, MOE} + VALID_LAYERS = {MAMBA, GDN, ATTENTION, DS_ATTENTION, CSA, HCA, WINDOW, MLA, MLP, MOE} # MLA-based attention layers (incompatible with standard '*' attention in one model). - MLA_ATTENTION = {DS_ATTENTION, CSA, HCA, WINDOW} + MLA_ATTENTION = {DS_ATTENTION, CSA, HCA, WINDOW, MLA} @classmethod def name_sorted_valid_layer_symbols(cls) -> list[str]: @@ -297,8 +298,8 @@ def _validate_pattern(pattern: str, pattern_name: str, allow_pipe: bool = False) f"Valid symbols are: {valid_chars}" ) - # Disallow standard Attention ('*') mixed with any MLA-based attention (D/C/H/W). - # MLA variants (DSA / CSA / HCA / Window) may freely coexist with each other. + # Disallow standard Attention ('*') mixed with any MLA-based attention (D/C/H/W/+). + # MLA variants (MLA / DSA / CSA / HCA / Window) may freely coexist with each other. if Symbols.ATTENTION in pattern and any(s in pattern for s in Symbols.MLA_ATTENTION): raise ValueError( "Not supported to have both Attention and MLA/DSA/CSA/HCA/Window in one model" @@ -328,7 +329,7 @@ def validate_segment_layers(segment: str) -> List[str]: f"one of {Symbols.VALID_LAYERS}" ) - # Disallow standard Attention ('*') mixed with any MLA-based attention (D/C/H/W). + # Disallow standard Attention ('*') mixed with any MLA-based attention (D/C/H/W/+). if Symbols.ATTENTION in segment and any(s in segment for s in Symbols.MLA_ATTENTION): raise ValueError( "Not supported to have both Attention and MLA/DSA/CSA/HCA/Window in one model" diff --git a/megatron/core/models/hybrid/hybrid_layer_specs.py b/megatron/core/models/hybrid/hybrid_layer_specs.py index 2f4e2cc1bf3..90e734fb465 100755 --- a/megatron/core/models/hybrid/hybrid_layer_specs.py +++ b/megatron/core/models/hybrid/hybrid_layer_specs.py @@ -174,6 +174,28 @@ self_attn_bda=get_bias_dropout_add, ), ), + mla_layer=ModuleSpec( + module=TransformerLayer, + submodules=TransformerLayerSubmodules( + input_layernorm=TENorm, + self_attention=ModuleSpec( + module=MLASelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=MLASelfAttentionSubmodules( + linear_q_proj=TEColumnParallelLinear, + linear_q_down_proj=TELinear, + linear_q_up_proj=TEColumnParallelLinear, + linear_kv_down_proj=TELinear, + linear_kv_up_proj=TEColumnParallelLinear, + core_attention=TEDotProductAttention, + linear_proj=TERowParallelLinear, + q_layernorm=IdentityOp, + kv_layernorm=IdentityOp, + ), + ), + self_attn_bda=get_bias_dropout_add, + ), + ), # Started with spec from gpt_layer_specs.py # Using the TE spec because we had problems getting the non-TE spec # working @@ -269,6 +291,28 @@ self_attn_bda=get_bias_dropout_add, ), ), + mla_layer=ModuleSpec( + module=TransformerLayer, + submodules=TransformerLayerSubmodules( + input_layernorm=TENorm, + self_attention=ModuleSpec( + module=MLASelfAttention, + params={"attn_mask_type": AttnMaskType.causal}, + submodules=MLASelfAttentionSubmodules( + linear_q_proj=TEColumnParallelLinear, + linear_q_down_proj=TELinear, + linear_q_up_proj=TEColumnParallelLinear, + linear_kv_down_proj=TELinear, + linear_kv_up_proj=TEColumnParallelLinear, + core_attention=TEDotProductAttention, + linear_proj=InferenceRowParallelLinear, + q_layernorm=IdentityOp, + kv_layernorm=IdentityOp, + ), + ), + self_attn_bda=get_bias_dropout_add, + ), + ), # Started with spec from gpt_layer_specs.py # Using the TE spec because we had problems getting the non-TE spec # working diff --git a/megatron/core/models/mimo/model/base.py b/megatron/core/models/mimo/model/base.py index 7485df0787f..b226b5c1e4b 100644 --- a/megatron/core/models/mimo/model/base.py +++ b/megatron/core/models/mimo/model/base.py @@ -2,6 +2,7 @@ import logging import warnings +from contextlib import ExitStack, contextmanager from typing import Any, Dict, Optional, Tuple import torch @@ -292,6 +293,51 @@ def _active_submodules(self): if submodule is not None: yield submodule + def _active_ddp_modules(self): + """Yield this rank's active DDP-wrapped submodules.""" + for module in self._active_submodules(): + if isinstance(module, DistributedDataParallel): + yield module + + @contextmanager + def no_sync(self): + """Disable grad-ready registration on overlapped inner DDP modules.""" + with ExitStack() as stack: + for module in self._active_ddp_modules(): + if module.ddp_config.overlap_grad_reduce: + stack.enter_context(module.no_sync()) + yield + + def enable_forward_pre_hook(self): + """Enable parameter-gather pre-hooks on overlapped inner DDP modules.""" + for module in self._active_ddp_modules(): + if module.ddp_config.overlap_param_gather: + module.enable_forward_pre_hook() + + def disable_forward_pre_hook(self, param_sync: bool = True): + """Disable parameter-gather pre-hooks on overlapped inner DDP modules.""" + for module in self._active_ddp_modules(): + if module.ddp_config.overlap_param_gather: + module.disable_forward_pre_hook(param_sync=param_sync) + + def start_param_sync(self, *unused, force_sync: bool = False, force_dispatch: bool = False): + """Start parameter synchronization on overlapped inner DDP modules.""" + for module in self._active_ddp_modules(): + if module.ddp_config.overlap_param_gather: + module.start_param_sync(force_sync=force_sync, force_dispatch=force_dispatch) + + def start_grad_sync(self, *unused): + """Start gradient synchronization on overlapped inner DDP modules.""" + for module in self._active_ddp_modules(): + if module.ddp_config.overlap_grad_reduce: + module.start_grad_sync() + + def free_overlap_buffers(self): + """Release parameter-gather buffers owned by overlapped inner DDP modules.""" + for module in self._active_ddp_modules(): + if module.ddp_config.overlap_param_gather: + module.free_overlap_buffers() + def zero_grad_buffer(self): """Zero each active submodule's DDP grad buffer.""" for module in self._active_submodules(): diff --git a/megatron/core/models/mimo/optimizer.py b/megatron/core/models/mimo/optimizer.py index 71500b5fcb6..598f6c883af 100644 --- a/megatron/core/models/mimo/optimizer.py +++ b/megatron/core/models/mimo/optimizer.py @@ -103,6 +103,11 @@ def step(self) -> Tuple[bool, Optional[float], Optional[int]]: num_zeros = self.count_zeros() if self.config.log_num_zeros_in_grad else None success = self.step_with_ready_grads() + # Reduce update success across the world (MIN) so disjoint-grid ranks agree. + success_tensor = torch.tensor([1 if success else 0], dtype=torch.int, device="cuda") + torch.distributed.all_reduce(success_tensor, op=torch.distributed.ReduceOp.MIN) + success = bool(success_tensor.item()) + return success, grad_norm, num_zeros @torch.no_grad() @@ -118,6 +123,11 @@ def zero_grad(self, set_to_none: bool = True): for opt in self._active_optimizers: opt.zero_grad(set_to_none) + def prepare_model_params_for_param_sync(self) -> None: + """Stage parameters for explicit synchronization in all active module optimizers.""" + for opt in self._active_optimizers: + opt.prepare_model_params_for_param_sync() + def get_loss_scale(self) -> torch.Tensor: """Return the loss scale tensor from the first active optimizer.""" if self._active_optimizers: diff --git a/megatron/core/models/mimo/submodules/base.py b/megatron/core/models/mimo/submodules/base.py index f05ecc6b15c..ac7bf64c063 100644 --- a/megatron/core/models/mimo/submodules/base.py +++ b/megatron/core/models/mimo/submodules/base.py @@ -310,13 +310,18 @@ def forward( Dictionary containing encoder-specific inputs. Keys should match encoder names. Used when is_first_stage=True. hidden_states (Optional[torch.Tensor]): - Hidden states from previous pipeline stage. Used when is_first_stage=False. + Already-combined encoder features. When supplied, bypasses encoding on any stage. Returns: Optional[torch.Tensor]: Processed and projected embeddings tensor, or None if no embeddings were produced. """ - if self.is_first_stage: + if encoder_inputs is not None and hidden_states is not None: + raise ValueError("encoder_inputs and hidden_states are mutually exclusive") + + if hidden_states is not None: + combined = hidden_states + elif self.is_first_stage: if encoder_inputs is None: return None embeddings = self.encode(encoder_inputs) @@ -324,9 +329,7 @@ def forward( return None combined = self.combine_embeddings(embeddings) else: - if hidden_states is None: - return None - combined = hidden_states + return None if self.is_last_stage: return self.project_embeddings([combined], is_input=True) diff --git a/megatron/core/msc_utils.py b/megatron/core/msc_utils.py index ce7cb685e25..c7dd7dc07d6 100644 --- a/megatron/core/msc_utils.py +++ b/megatron/core/msc_utils.py @@ -55,13 +55,57 @@ def __setstate__(self, state): MultiStorageClientFeature = _FeatureFlag(default=False) -def open_file(*args, **kwargs): - """Open a file with the appropriate method based on whether MSC is enabled.""" - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - return msc.open(*args, **kwargs) - else: - return open(*args, **kwargs) - - -__all__ = ['MultiStorageClientFeature', 'open_file'] +class MaybeMultiStorageClient: + """ + Helper class to use MultiStorageClient + """ + + def path_isdir(self, path, strict: bool = True): + """ + Check if a path is an existing directory. + :param path: path to check + :param strict: if True, use only committed metadata for MSC + """ + if MultiStorageClientFeature.is_enabled(): + pkg = MultiStorageClientFeature.import_package() + return pkg.os.path.isdir(path, strict=strict) + else: + import os + + return os.path.isdir(path) + + def __getattr__(self, name): + if MultiStorageClientFeature.is_enabled(): + pkg = MultiStorageClientFeature.import_package() + if hasattr(pkg, name): + return getattr(pkg, name) + + if name == "open": + return open + if name == "os": + import os + + return os + if name == "Path": + from pathlib import Path + + return Path + if name == "torch": + import torch + + return torch + raise AttributeError(f"{self.__class__.__name__!s} has no attribute {name!s}") + + def __dir__(self): + attrs = {"open", "os", "Path", "torch"} + if MultiStorageClientFeature.is_enabled(): + try: + pkg = MultiStorageClientFeature.import_package() + attrs.update(dir(pkg)) + except RuntimeError: + pass + return sorted(attrs) + + +maybe_msc = MaybeMultiStorageClient() +__all__ = ['MultiStorageClientFeature', 'maybe_msc'] diff --git a/megatron/core/optimizer/__init__.py b/megatron/core/optimizer/__init__.py index 27b675d1b8d..70f757f2889 100644 --- a/megatron/core/optimizer/__init__.py +++ b/megatron/core/optimizer/__init__.py @@ -47,7 +47,6 @@ if HAVE_EMERGING_OPTIMIZERS: from emerging_optimizers.scalar_optimizers import Lion -from megatron.core import parallel_state from megatron.core.optimizer.cpu_offloading.hybrid_optimizer import HybridDeviceOptimizer from megatron.core.optimizer_param_scheduler import ( ParamGroupOverride, @@ -686,11 +685,12 @@ def init_state_fn(opt, config=None): setattr(optimizer, 'grad_stats_parallel_group', model_parallel_group) if pg_collection is None or not hasattr(pg_collection, 'tp'): - tp_group = parallel_state.get_tensor_model_parallel_group() - else: - tp_group = pg_collection.tp - # TODO(M4): plumb tp_group through optimizer constructors so this setattr disappears. + pg_collection = ProcessGroupCollection.use_mpu_process_groups() + tp_group = pg_collection.tp + expert_tp_group = getattr(pg_collection, 'expt_tp', tp_group) + # TODO(M4): plumb TP groups through optimizer constructors so these setattrs disappear. setattr(optimizer, 'tp_group', tp_group) + setattr(optimizer, 'expert_tp_group', expert_tp_group) return optimizer @@ -798,7 +798,15 @@ def _get_megatron_emerging_optimizer( ) # Apply optimizer-specific default param overrides (e.g. muon: non-linear -> adam). - config_overrides.update(_EMERGING_OPTIMIZERS[eopt_name].default_param_overrides) + # For Muon-family optimizers, the scalar optimizer that handles non-linear/embedding + # params is configurable via ``config.muon_scalar_optimizer`` (e.g., 'adam' or 'lion'); + # deep-copy the registry defaults before rewriting so we never mutate shared state. + default_param_overrides = copy.deepcopy(_EMERGING_OPTIMIZERS[eopt_name].default_param_overrides) + if eopt_name in ('muon', 'adaptive_muon'): + for override in default_param_overrides.values(): + if override.get('optimizer') in ('adam', 'lion'): + override['optimizer'] = config.muon_scalar_optimizer + config_overrides.update(default_param_overrides) # Build param groups and bucket by (optimizer_name, is_expert_parallel). # Layer-wise distributed optimizer handles expert params internally so we skip that split. @@ -831,7 +839,10 @@ def _get_megatron_emerging_optimizer( "fall back to the legacy LayerWise ping-pong path." ) if use_separate_distributed_optimizer and any( - opt_name not in _EMERGING_OPTIMIZERS + # A separate DistributedOptimizer with byte-level sharding handles any group + # whose optimizer is not the primary emerging optimizer (stored in ``eopt_name``, + # e.g., Muon). This includes scalar optimizers like Adam or Lion. + not (opt_name == eopt_name and opt_name in _EMERGING_OPTIMIZERS) for (opt_name, _), groups in grouped_param_groups.items() if groups ): @@ -870,7 +881,10 @@ def _get_megatron_emerging_optimizer( model_parallel_group = pg_collection.tp_ep_pp if is_expert else pg_collection.mp - if opt_name in _EMERGING_OPTIMIZERS: + # Only the primary emerging optimizer (stored in ``eopt_name``, e.g., Muon) is + # constructed via ``_create_emerging_optimizer``. Scalar optimizers that also appear + # in ``_EMERGING_OPTIMIZERS`` (e.g., Lion) fall through to the standard fallback path. + if opt_name == eopt_name and opt_name in _EMERGING_OPTIMIZERS: optimizer, init_state_fn = _create_emerging_optimizer( config, groups, eopt_name, model_chunks, pg_collection ) @@ -884,18 +898,17 @@ def _get_megatron_emerging_optimizer( else: optimizer = FP32Optimizer(optimizer, config, init_state_fn) setattr(optimizer, 'grad_stats_parallel_group', model_parallel_group) - if pg_collection is None or not hasattr(pg_collection, 'tp'): - tp_group = parallel_state.get_tensor_model_parallel_group() - else: - tp_group = pg_collection.tp + tp_group = pg_collection.tp + expert_tp_group = getattr(pg_collection, 'expt_tp', tp_group) setattr(optimizer, 'tp_group', tp_group) + setattr(optimizer, 'expert_tp_group', expert_tp_group) results.append(optimizer) continue else: fallback_config = copy.copy(config) fallback_config.optimizer = opt_name if use_separate_distributed_optimizer: - # Route non-emerging params through a real DistributedOptimizer + # Route non-emerging params (adam/lion) through a real DistributedOptimizer # (byte-level sharding) instead of stuffing them inside LayerWise. for group in groups: assert not group['is_expert_parallel'], ( @@ -1050,10 +1063,12 @@ def get_megatron_optimizer( intra_expt_dp_group = process_groups_dict['intra_expt_dp_group'] mp_group = process_groups_dict['mp_group'] expt_tp_pp_group = process_groups_dict['expt_tp_pp_group'] + expt_tp_pp_with_egtp_remat_group = process_groups_dict['expt_tp_pp_with_egtp_remat_group'] intra_dp_cp_group_gloo = process_groups_dict['intra_dp_cp_group_gloo'] intra_expt_dp_group_gloo = process_groups_dict['intra_expt_dp_group_gloo'] intra_dist_opt_group = process_groups_dict['intra_dist_opt_group'] + # ``mp_group`` spans TP×GTP_remat×PP (GTP_remat-merged). model_parallel_rank = get_pg_rank(mp_group) if get_pg_size(dp_cp_group) > get_pg_size(intra_dp_cp_group): @@ -1178,8 +1193,9 @@ def get_megatron_optimizer( param_to_param_group[param_name] = param_group_id param_group_id += 1 if len(moe_param_groups) > 0: - expt_model_parallel_rank = get_pg_rank(expt_tp_pp_group) - # Pass Gloo process groups into optimizer only if needed. + # Expert analog of dense ``model_parallel_rank``; use the EGTP_remat-merged group so each + # EGTP_remat peer gets a distinct distopt ShardedObject key (else DCP "duplicate" error). + expt_model_parallel_rank = get_pg_rank(expt_tp_pp_with_egtp_remat_group) if use_gloo_process_groups: expt_data_parallel_group_gloo = intra_expt_dp_group_gloo else: @@ -1190,7 +1206,7 @@ def get_megatron_optimizer( model_chunks=model_chunks, param_groups=moe_param_groups, per_model_buffers=moe_buffers, - model_parallel_group=expt_tp_pp_group, + model_parallel_group=expt_tp_pp_with_egtp_remat_group, data_parallel_group=intra_expt_dp_group, data_parallel_group_gloo=expt_data_parallel_group_gloo, data_parallel_group_idx=expt_model_parallel_rank, diff --git a/megatron/core/optimizer/clip_grads.py b/megatron/core/optimizer/clip_grads.py index 3c5491d39a1..ad2064ede36 100644 --- a/megatron/core/optimizer/clip_grads.py +++ b/megatron/core/optimizer/clip_grads.py @@ -47,7 +47,7 @@ multi_tensor_scale_tensor_impl = None -from ..tensor_parallel import param_is_not_tensor_parallel_duplicate +from ..tensor_parallel import param_is_not_gtp_duplicate, param_is_not_tensor_parallel_duplicate from ..transformer.module import param_is_not_shared from ..utils import get_data_parallel_group_if_dtensor, to_local_if_dtensor @@ -92,7 +92,7 @@ def get_grad_norm_fp32( # Calculate norm. if norm_type == inf: - total_norm = max(grad.abs().max() for grad in grads_for_norm) + total_norm = max((grad.abs().max() for grad in grads_for_norm), default=torch.tensor(0.0)) total_norm_cuda = torch.tensor([float(total_norm)], dtype=torch.float, device='cuda') # Take max across all data-parallel GPUs if using FSDP and then all model-parallel GPUs. if data_parallel_group: @@ -105,24 +105,20 @@ def get_grad_norm_fp32( total_norm = total_norm_cuda[0].item() else: - if norm_type == 2.0: + total_norm = torch.zeros(1, dtype=torch.float, device='cuda') + if not grads_for_norm: + pass + elif norm_type == 2.0: dummy_overflow_buf = torch.zeros(1, dtype=torch.int, device='cuda') # Use apex's multi-tensor applier for efficiency reasons. # Multi-tensor applier takes a function and a list of list # and performs the operation on that list all in one kernel. - if grads_for_norm: - grad_norm, _ = multi_tensor_applier( - l2_norm_impl, - dummy_overflow_buf, - [grads_for_norm], - False, # no per-parameter norm - ) - else: - grad_norm = torch.zeros(1, dtype=torch.float, device='cuda') + grad_norm, _ = multi_tensor_applier( + l2_norm_impl, dummy_overflow_buf, [grads_for_norm], False # no per-parameter norm + ) # Since we will be summing across data parallel groups, # we need the pow(norm-type). total_norm = grad_norm**norm_type - else: for grad in grads_for_norm: grad_norm = torch.norm(grad, norm_type) @@ -201,14 +197,15 @@ def count_zeros_fp32( grad_stats_parallel_group: torch.distributed.ProcessGroup, use_decoupled_grad: bool = False, tp_group: Optional[torch.distributed.ProcessGroup] = None, + expert_tp_group: Optional[torch.distributed.ProcessGroup] = None, ) -> float: """Counts the number of zero values in the gradients of the given parameters. The count is performed in FP32. This method filters parameters to ensure gradients are not double-counted by checking if the gradient is not None, - the parameter is not shared, and the parameter is not a replica due - to tensor model parallelism. It also handles parameters managed by - Megatron FSDP specifically. + the parameter is not shared, and the parameter is not a replica due to + tensor model parallelism or (expert) generalized tensor parallelism. It also + handles parameters managed by Megatron FSDP specifically. Args: parameters (Union[List[torch.Tensor], torch.Tensor]): An iterable of @@ -230,6 +227,7 @@ def count_zeros_fp32( # - grad should not be none # - parameter should not be shared # - should not be a replica due to tensor model parallelism + # - should not be a replica due to (expert) generalized tensor parallelism total_num_zeros = torch.zeros(1, dtype=torch.int64, device='cuda') data_parallel_group = None use_megatron_fsdp = False @@ -245,8 +243,11 @@ def count_zeros_fp32( total_num_zeros += num_zeros continue is_not_shared = param_is_not_shared(param) - is_not_tp_duplicate = param_is_not_tensor_parallel_duplicate(param, tp_group=tp_group) - if grad_not_none and is_not_shared and is_not_tp_duplicate: + is_not_tp_duplicate = param_is_not_tensor_parallel_duplicate( + param, tp_group=tp_group, expert_tp_group=expert_tp_group + ) + is_not_gtp_duplicate = param_is_not_gtp_duplicate(param) + if grad_not_none and is_not_shared and is_not_tp_duplicate and is_not_gtp_duplicate: grad_obj = getattr(param, grad_attr) data_parallel_group = get_data_parallel_group_if_dtensor(grad_obj, data_parallel_group) grad = to_local_if_dtensor(grad_obj).detach() diff --git a/megatron/core/optimizer/distrib_optimizer.py b/megatron/core/optimizer/distrib_optimizer.py index 9baed75a11d..5f50db187ac 100644 --- a/megatron/core/optimizer/distrib_optimizer.py +++ b/megatron/core/optimizer/distrib_optimizer.py @@ -212,6 +212,18 @@ def _build_model_gbuf_range(cls, param_and_grad_buffer: _ParamAndGradBuffer, buc data_parallel_rank = param_and_grad_buffer.data_parallel_group.rank() data_parallel_world_size = param_and_grad_buffer.data_parallel_group.size() + # The layout records how many shards it was built for. That count has to match the group + # the reduce-scatter and all-gather run over, which is the intra-instance group when + # there are several optimizer instances. If the layout was sized by a larger group, the + # trailing shards of every bucket belong to no rank: those params are never updated and + # drop out of grad-norm, num-zeros and params-norm, which sum over owned shards only. + num_optimizer_shards = param_and_grad_buffer.num_optimizer_shards + assert num_optimizer_shards is None or num_optimizer_shards == data_parallel_world_size, ( + f"Parameter layout was built for {num_optimizer_shards} optimizer shards but the " + f"buffer's data-parallel group has {data_parallel_world_size} ranks. Size the layout " + f"by the group the optimizer shards over." + ) + bucket = param_and_grad_buffer.buckets[bucket_index] gbuf_size = bucket.grad_data.numel() assert ( @@ -424,6 +436,7 @@ def _build_model_and_main_param_groups( tensor_parallel.copy_tensor_model_parallel_attributes( shard_model_param, model_param ) + tensor_parallel.copy_gtp_attributes(shard_model_param, model_param) copy_optimizer_param_metadata(shard_model_param, model_param) # Generate main param. @@ -455,6 +468,7 @@ def _build_model_and_main_param_groups( tensor_parallel.copy_tensor_model_parallel_attributes( shard_main_param, model_param ) + tensor_parallel.copy_gtp_attributes(shard_main_param, model_param) copy_optimizer_param_metadata(shard_main_param, model_param) else: # When using precision-aware optimizer, main params are held by FusedAdam. @@ -477,6 +491,7 @@ def _build_model_and_main_param_groups( tensor_parallel.copy_tensor_model_parallel_attributes( shard_model_param, model_param ) + tensor_parallel.copy_gtp_attributes(shard_model_param, model_param) copy_optimizer_param_metadata(shard_model_param, model_param) else: @@ -590,6 +605,7 @@ def _finalize_bucket(param_end_index, bucket_start_index, bucket_id): bucket_indices=bucket_indices, per_bucket_numel_unpadded=per_bucket_numel_unpadded, param_indices=param_indices if param_indices is not None else [], + num_optimizer_shards=data_parallel_world_size, ) @staticmethod @@ -697,9 +713,10 @@ def __init__( assert ( isinstance(optimizer, (Adam, torch.optim.AdamW, HybridDeviceOptimizer)) or optimizer is None + or init_state_fn is not None ), ( - "Only Adam and HybridDeviceOptimizer currently supported, " - "due to checkpointing requirements." + "Only Adam, HybridDeviceOptimizer, and optimizers with an init_state_fn " + "(e.g., Lion) are currently supported, due to checkpointing requirements." ) self._state_offloader: Optional[OptimizerStateOffloader] = None @@ -833,6 +850,27 @@ def get_grad_stats_parallel_group(self) -> torch.distributed.ProcessGroup: """ return getattr(self, 'grad_stats_parallel_group', None) + @property + def optimizer_state_keys(self): + """Return the optimizer's tensor state keys, e.g., ('exp_avg', 'exp_avg_sq') for Adam + or ('exp_avg',) for Lion.""" + _OPTIMIZER_STATE_KEYS = {"lion": ("exp_avg",)} + optimizer_name = self.config.optimizer + # When Muon is the top-level optimizer, the DistributedOptimizer wrapping + # scalar parameters uses muon_scalar_optimizer (e.g., Lion) as the actual + # optimizer, so look up state keys by that name instead. + if optimizer_name == "muon": + optimizer_name = self.config.muon_scalar_optimizer + return _OPTIMIZER_STATE_KEYS.get(optimizer_name, ("exp_avg", "exp_avg_sq")) + + def _get_state_key_dtype(self, key): + """Return the dtype for a given optimizer state key.""" + dtype_map = { + "exp_avg": self.config.exp_avg_dtype, + "exp_avg_sq": self.config.exp_avg_sq_dtype, + } + return dtype_map.get(key, torch.float32) + def state_dict(self): """ The state dict contains all non-DP-rank-dependent (i.e., non-parameter- @@ -944,23 +982,31 @@ def load_state_dict(self, state_dict): # contains an integer ordering of parameters within each group, and # the ordering of parameters within its flattened parameter state # list. + + # Pair each current param_group with its saved counterpart by identifier tuple. + # Construction order isn't part of the checkpoint, so we match by a tuple of + # per-group config (``param_group_identifier_keys``) rather than by position. + def make_needed_groups(param_group): needed_groups = [] for key in param_group_identifier_keys: - # NeMo changes these variable names from `lr_mult` and `wd_mult` - # to `pre_lr_mult` and `pre_wd_mult`, so we need to check both. + # NeMo aliases ``lr_mult``/``wd_mult`` as ``pre_lr_mult``/``pre_wd_mult``. if key in param_group: - pass + value = param_group[key] elif f"pre_{key}" in param_group: - key = f"pre_{key}" + value = param_group[f"pre_{key}"] else: - raise ValueError( - f"Key {key} (or pre_{key}) not found in param_group {param_group}." - ) - needed_groups.append(param_group[key]) - needed_groups = tuple(needed_groups) - return needed_groups - + # Treat missing and explicit None identifier values as equivalent. + value = None + needed_groups.append(value) + return tuple(needed_groups) + + # Duplicate identifiers here silently clobber: two saved groups with the same tuple + # collapse to whichever was inserted last, and one current group inherits the wrong + # override state (``max_lr`` etc.). Params are unaffected — they come from the + # inner optimizer below — but the next step runs at the wrong LR / WD. Adding the + # distinguishing field to ``param_group_identifier_keys`` is the fix. See + # ``test_filter_reorder_distinguishes_groups_by_max_lr``. param_groups_map = {} for param_group in state_dict["optimizer"]["param_groups"]: needed_groups = make_needed_groups(param_group) @@ -1004,8 +1050,8 @@ def make_needed_groups(param_group): # For precision_aware_optimizer, the empty tensors should also be # initialized with the correct dtype. tensors = { - "exp_avg": init_shard(self.config.exp_avg_dtype), - "exp_avg_sq": init_shard(self.config.exp_avg_sq_dtype), + key: init_shard(self._get_state_key_dtype(key)) + for key in self.optimizer_state_keys } if self.config.use_precision_aware_optimizer_no_fp8_or_ds_fp8: if self.config.store_param_remainders and self.config.bf16: @@ -1094,12 +1140,8 @@ def _get_main_param_and_optimizer_states(self, model_param): """Return a dict containing the main param and optimizer states corresponding to the input model_param. - The structure of the returned dict: - tensors = { - "param": torch.Tensor - "exp_avg": torch.Tensor - "exp_avg_sq": torch.Tensor - } + The returned dict always contains "param" and one entry per optimizer state tensor + (e.g., "exp_avg" and "exp_avg_sq" for Adam, or just "exp_avg" for Lion). """ group_index, group_order = self.model_param_group_index_map[model_param] if self.config.use_precision_aware_optimizer_no_fp8_or_ds_fp8: @@ -1200,12 +1242,8 @@ def _expand_quantized_param_shard_for_cast( def _set_main_param_and_optimizer_states(self, model_param, tensors): """Set the main param and optimizer states corresponding to the input model_param. - The structure of the input `tensors`: - tensors = { - "param": torch.Tensor - "exp_avg": torch.Tensor - "exp_avg_sq": torch.Tensor - } + The input `tensors` dict contains "param" and one entry per optimizer state tensor + (e.g., "exp_avg" and "exp_avg_sq" for Adam, or just "exp_avg" for Lion). """ group_index, group_order = self.model_param_group_index_map[model_param] if self.config.use_precision_aware_optimizer_no_fp8_or_ds_fp8: @@ -1333,7 +1371,7 @@ def get_parameter_state_dp_zero( key: torch.zeros( (buffer_numel_unpadded,), dtype=torch.float32, device="cpu" ) - for key in ("param", "exp_avg", "exp_avg_sq") + for key in ("param",) + self.optimizer_state_keys } world_tensors["numel_unpadded"] = buffer_numel_unpadded @@ -1355,7 +1393,7 @@ def get_parameter_state_dp_zero( local_shards = { key: torch.zeros((gbuf_local_numel,), dtype=torch.float32, device="cpu") - for key in ("param", "exp_avg", "exp_avg_sq") + for key in ("param",) + self.optimizer_state_keys } # Build contiguous DP rank shards (for param + optim states). @@ -1714,8 +1752,8 @@ def sharded_param_state_fully_reshardable( `fully_reshardable` format involves gathering the tensors on DP rank 0 during save. Flat DistOpt buffers are unflattened and reshaped into model param like sizes. This results in a state dict similar to a regular optimizer one, where each - param of shape (X, Y, Z) has corresponding 'param', 'exp_avg' and 'exp_avg_sq' - tensors of shape (X, Y, Z) in the optimizer state dict. + param of shape (X, Y, Z) has corresponding 'param' and optimizer state + tensors (e.g., 'exp_avg', 'exp_avg_sq') of shape (X, Y, Z) in the optimizer state dict. During loading there is no data exchange - each rank requests to load the whole state dict (and flattens and trims the tensors afterwards). It is recommended @@ -2187,7 +2225,7 @@ def load_parameter_state_from_dp_zero_legacy(self, state_dict): t.numel() for t in state_dict[gbuf_idx][torch.float32]["param"] ] assert sum(model_numels) == sum(checkpoint_numels) - for key in ("param", "exp_avg", "exp_avg_sq"): + for key in ("param",) + self.optimizer_state_keys: legacy_world_tensors = self._update_legacy_world_tensors( state_dict[gbuf_idx][torch.float32][key], [ @@ -2308,7 +2346,7 @@ def load_parameter_state_from_dp_zero(self, state_dict, *, update_legacy_format= f"({buffer_numel_unpadded}) and checkpoint ({checkpoint_numel_unpadded})" ) recv_tensors = {} - for key in ("param", "exp_avg", "exp_avg_sq"): + for key in ("param",) + self.optimizer_state_keys: offset_in_world_tensors = 0 for bucket_idx, gbuf_range_map in enumerate(gbuf_range_map_for_all_buckets): # Compute local DP contiguous shard's size. @@ -2521,7 +2559,7 @@ def split_state_dict_if_needed(self, state_dict): # Split the target buffer into two separate buffers. fp8_state_dict, non_fp8_state_dict = {}, {} - for key in ['param', 'exp_avg', 'exp_avg_sq']: + for key in ('param',) + self.optimizer_state_keys: tensor = state_dict[non_fp8_gbuf_idx][non_fp8_param_and_grad_dtype][key] fp8_tensor = torch.empty([fp8_offsets[-1]], dtype=tensor.dtype) non_fp8_tensor = torch.empty([non_fp8_offsets[-1]], dtype=tensor.dtype) diff --git a/megatron/core/optimizer/emerging_optimizers.py b/megatron/core/optimizer/emerging_optimizers.py index 8d3197708ea..7e936f719eb 100644 --- a/megatron/core/optimizer/emerging_optimizers.py +++ b/megatron/core/optimizer/emerging_optimizers.py @@ -17,7 +17,7 @@ from torch.optim.optimizer import ParamsT from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.utils import get_pg_size, log_single_rank +from megatron.core.utils import get_pg_rank, get_pg_size, log_single_rank from .optimizer_config import ParamKey, ParamPredicate @@ -230,6 +230,42 @@ def scaled_orthogonalize_fn( scaled_orthogonalize_fn=scaled_orthogonalize_fn, ) + def scaled_orthogonalize_fn_with_gtp_remat(self, p, grad, tp_group, partition_dim): + """All-gather grad along GTP_remat/EGTP_remat dim 0, orthogonalize, then slice back. + + GTP_remat shards weights along dim 0 independently of TP's partition_dim. Newton-Schulz + needs the full weight matrix, so we reconstruct the GTP_remat dimension before running + the TP-aware orthogonalization, then extract the local GTP_remat shard from the result. + When GTP_remat is inactive this is a plain passthrough to scaled_orthogonalize_fn. + """ + # TODO: Clean up code that determines if parameter is a MoE layer and which TP group to use + is_expert = getattr(p, 'expert_tp', False) + gtp_remat_group = ( + (self.pg_collection.expt_gtp_remat if is_expert else self.pg_collection.gtp_remat) + if self.pg_collection + else None + ) + + if gtp_remat_group is None or get_pg_size(gtp_remat_group) <= 1: + return self.scaled_orthogonalize_fn(grad, tp_group, partition_dim) + + # Parameters with is_gtp_weight_remat=False are not sharded along the + # GTP process group, and do not require all-gathering prior to + # orthogonalization. + if not getattr(p, 'is_gtp_weight_remat', False): + return self.scaled_orthogonalize_fn(grad, tp_group, partition_dim) + + gtp_remat_size = get_pg_size(gtp_remat_group) + gtp_rank = get_pg_rank(gtp_remat_group) + shards = [torch.empty_like(grad) for _ in range(gtp_remat_size)] + torch.distributed.all_gather(shards, grad, gtp_remat_group) + gathered_grad = torch.cat(shards, dim=0) + + gathered_grad = self.scaled_orthogonalize_fn(gathered_grad, tp_group, partition_dim) + + shard_size = gathered_grad.shape[0] // gtp_remat_size + return gathered_grad[gtp_rank * shard_size : (gtp_rank + 1) * shard_size].contiguous() + def orthogonalize(self, p: torch.Tensor, grad: torch.Tensor, **kwargs: Any) -> torch.Tensor: """Orthogonalize the momentum. @@ -280,14 +316,14 @@ def orthogonalize(self, p: torch.Tensor, grad: torch.Tensor, **kwargs: Any) -> t qkv_grads = [g.reshape(-1, grad_shape[-1]) for g in qkv_grads] qkv_grads = [ - self.scaled_orthogonalize_fn(g, tp_group, partition_dim).view( + self.scaled_orthogonalize_fn_with_gtp_remat(p, g, tp_group, partition_dim).view( num_query_groups, -1, grad_shape[-1] ) for g in qkv_grads ] grad = torch.cat(qkv_grads, dim=1).view(grad_shape) else: - grad = self.scaled_orthogonalize_fn(grad, tp_group, partition_dim) + grad = self.scaled_orthogonalize_fn_with_gtp_remat(p, grad, tp_group, partition_dim) return grad diff --git a/megatron/core/optimizer/layer_wise_optimizer.py b/megatron/core/optimizer/layer_wise_optimizer.py index bb12d159437..376f9a1f1c0 100644 --- a/megatron/core/optimizer/layer_wise_optimizer.py +++ b/megatron/core/optimizer/layer_wise_optimizer.py @@ -1,7 +1,8 @@ -# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import logging import math +import re from typing import Callable, Dict, List, Optional, Tuple import torch @@ -13,12 +14,6 @@ from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.utils import get_pg_rank, get_pg_size, log_single_rank -from ..fp8_utils import ( - _stage_param_to_bf16, - copy_back_gathered_bf16_into_fp8_param, - is_float8tensor, - post_all_gather_processing, -) from .clip_grads import count_zeros_fp32, get_grad_norm_fp32 from .optimizer import ( ChainedOptimizer, @@ -76,16 +71,6 @@ def _bucket_is_managed_by_layer_wise_optimizer(bucket, default_for_untagged: boo return param.is_managed_by_layer_wise_optimizer -def _param_sort_key(numel: int, identity: tuple) -> tuple: - """Rank-independent total-order key for ping-pong ownership: ``(numel, *canonical-identity)``. - - ``numel`` alone is not a total order (stable sort tie-breaks by rank-local insertion order), - so equal-numel params would get different owners across ranks; the canonical identity - ``(chunk_idx, buffer_idx, global_start_index)`` makes it total and identical on every rank. - """ - return (numel,) + tuple(identity) - - def tag_params_for_buffer_routing(model_chunks) -> None: """Tag every requires-grad param with ``is_managed_by_layer_wise_optimizer``. @@ -102,6 +87,83 @@ def tag_params_for_buffer_routing(model_chunks) -> None: param.is_managed_by_layer_wise_optimizer = is_managed_by_layer_wise_optimizer(param) +def _build_gtp_replica_fold(pg_collection, model_chunks) -> Dict[str, Tuple[int, int]]: + """Map each (E)GTP_remat-REPLICATED param to ``(gtp_rank, gtp_remat_size)`` for folding. + + PROBLEM: LayerWise keeps (E)GTP_remat-replicated params (identical per gtp_remat peer) WHOLE, so + their optimizer-state ShardedTensors share one key+offset across those peers. The DP-coord reset + in ``sharded_state_dict`` would then mark all peers the all-zero "main replica" -> DCP sees N + writers for one shard and rejects the save. + + FIX: fold the (e)gtp_remat rank into ``replica_id[1]`` so one peer writes. (E)GTP_remat-SHARDED + params (``is_gtp_param``) are offset-sharded and excluded -- each shard already has a + distinct offset, hence a unique writer. + + Returns: ``{param_name: (gtp_rank, gtp_remat_size)}``, empty when GTP_remat is unavailable or + no group spans >1 rank. Names are bare (all ``module.`` wrappers stripped, layer index + collapsed) to match the optimizer-state checkpoint key suffix. + """ + gtp_fold: Dict[str, Tuple[int, int]] = {} + try: + from megatron.core.tensor_parallel.gtp_api import HAVE_GTP, is_gtp_param + except ImportError: + return gtp_fold + if not HAVE_GTP: + return gtp_fold + + assert pg_collection is not None, ( + "_build_gtp_replica_fold requires a pg_collection carrying gtp_remat/expt_gtp_remat; " + "the optimizer factory must materialize it before constructing the optimizer." + ) + gtp_remat_group = getattr(pg_collection, 'gtp_remat', None) + egtp_remat_group = getattr(pg_collection, 'expt_gtp_remat', None) + + for model_chunk in model_chunks: + for name, p in model_chunk.named_parameters(): + if is_gtp_param(p): + continue + grp = egtp_remat_group if getattr(p, 'is_expert_parallel', False) else gtp_remat_group + if grp is None or grp.size() <= 1: + continue + # Normalize the param name so it matches the optimizer-state checkpoint key suffix, + # which is wrapper-free and layer-collapsed. Three transforms, in order: + # 1. drop every leading 'module.' (DDP + Float16Module can double-wrap the model), + # 2. collapse the layer index (the checkpoint key drops it -- a sharded axis), and + # 3. collapse SequentialMLP 'local_experts.' to the grouped key 'experts' (the + # checkpoint groups them, matching TEGroupedMLP), else expert replicas collide. + # e.g. 'module.module.decoder.layers.3.mlp.router.weight' + # -> 'decoder.layers.mlp.router.weight' + nm = name + while nm.startswith('module.'): + nm = nm[len('module.') :] + nm = re.sub(r'\.layers\.\d+\.', '.layers.', nm) + nm = re.sub(r'\.local_experts\.\d+\.', '.experts.', nm) + gtp_fold[nm] = (grp.rank(), grp.size()) + return gtp_fold + + +def _fold_replica_id(replica_id, key, gtp_fold: Dict[str, Tuple[int, int]]): + """Compute a ShardedTensor's writer-disambiguating replica_id for fixed-DP checkpointing. + + Base reset: keep (PP, TP), zero DP -- every DP rank holds the same shard, so one writer + remains. Correct for normal params. + + For an (e)gtp-replicated param (in ``gtp_fold``), reset leaves ``gtp_remat_size`` writers, so + fold the peer gtp_remat rank into TP slot to re-spread: ``new_tp = old_tp * gtp_remat_size + + gtp_rank`` (rank 0 stays the writer, the others move off the all-zero main replica) -> one + writer per shard. Suffix-match (bare fold name vs fully-qualified key) and collapse the key's + layer index too, so it matches per-layer and already-collapsed keys. + """ + rid = (*replica_id[:2], 0) + if not gtp_fold: + return rid + key = re.sub(r'\.layers\.\d+\.', '.layers.', key or '') + for nm, (gtp_rank, gtp_remat_size) in gtp_fold.items(): + if key.endswith(nm): + return (rid[0], rid[1] * gtp_remat_size + gtp_rank, rid[2]) + return rid + + class LayerWiseDistributedOptimizer(ChainedOptimizer): """Layer-wise distributed optimizer for Megatron-core models. @@ -304,6 +366,7 @@ def _emit_bucket( bucket_indices=bucket_indices, per_bucket_numel_unpadded=per_bucket_numel_unpadded, param_indices=param_indices if param_indices is not None else [], + num_optimizer_shards=dp_size, ) @staticmethod @@ -333,22 +396,9 @@ def compute_full_param_layout( :class:`FullParamLayout` with a :class:`PerBufferParamLayout` per buffer group. """ # Avoid a circular import: DistributedOptimizer imports LayerWise indirectly. - from ..distributed.param_and_grad_buffer import _compute_default_per_buffer_param_layout from .distrib_optimizer import DistributedOptimizer - # Decoupled layout (use_layer_wise_param_layout=False): LayerWise (Muon) buffers use a - # compact no-padding DDP layout (and locally disable DistributedOptimizer semantics in - # DDP), so they must NOT receive the shard-aligned ``dp_size * max(shard_load)`` padded - # layout here. Non-LayerWise buffers keep DistOpt's byte-level layout regardless. - decouple_ddp_layout = not ddp_config.use_layer_wise_param_layout - - # fp8 Muon grads key to uint8 (own buffer); partition_buckets later merges the non-fp8 - # bucket groups into the fp8 group to aggregate communication. - buffer_groups = group_params_for_buffers( - params, - ddp_config.grad_reduce_in_fp32, - merge_layerwise_fp8_grads=not getattr(ddp_config, 'use_layer_wise_param_layout', True), - ) + buffer_groups = group_params_for_buffers(params, ddp_config.grad_reduce_in_fp32) layouts = {} for buffer_key, (group_params, param_indices) in buffer_groups.items(): if buffer_key.is_expert_parallel: @@ -363,15 +413,6 @@ def compute_full_param_layout( # Dispatch per buffer: LayerWise (Muon) params get the shard-aligned # layout; non-LayerWise params (e.g. Adam-managed embeddings, biases) # get DistOpt's byte-level layout. - if buffer_key.is_managed_by_layer_wise_optimizer and decouple_ddp_layout: - # Decouple path (incl. FP8 param-gather): compact no-padding layout (DDP treats this - # buffer as non-DistOpt). Attach param_indices so DDP's consistency check passes. - per_buffer_layout = _compute_default_per_buffer_param_layout( - group_params, bucket_size - ) - per_buffer_layout.param_indices = param_indices - layouts[buffer_key] = per_buffer_layout - continue if buffer_key.is_managed_by_layer_wise_optimizer: compute_per_buffer_layout = ( LayerWiseDistributedOptimizer._compute_per_buffer_param_layout @@ -403,45 +444,43 @@ def __init__( """ self.pg_collection = pg_collection - self.decouple_ddp_layout = not config.use_layer_wise_param_layout + + # The data-parallel groups this optimizer shards parameters over. Cached here so the + # sharding, all-gather and broadcast paths read one attribute instead of reaching back + # into pg_collection at every use. + self.dp_cp = getattr(pg_collection, 'dp_cp', None) if pg_collection is not None else None + self.expt_dp = ( + getattr(pg_collection, 'expt_dp', None) if pg_collection is not None else None + ) + + # LayerWise assigns whole params to ranks of the full dp_cp group and all-gathers over + # that same group, so it has no notion of optimizer instances. With more than one + # instance, DDP reduce-scatters gradients over the smaller intra-instance group, so a + # rank would be asked to update params whose gradients it does not hold. Reject the + # combination instead of silently training on partial gradients. + intra_dp_cp = ( + getattr(pg_collection, 'intra_dp_cp', None) if pg_collection is not None else None + ) + assert intra_dp_cp is None or get_pg_size(intra_dp_cp) == get_pg_size(self.dp_cp), ( + "LayerWiseDistributedOptimizer does not support " + "num_distributed_optimizer_instances > 1." + ) full_param_layouts = None - if model_chunks is not None and not self.decouple_ddp_layout: + if model_chunks is not None: full_param_layouts = [ chunk.full_param_layout for chunk in model_chunks if hasattr(chunk, 'full_param_layout') and chunk.full_param_layout is not None ] or None - # Decouple path keeps whole-matrix ping-pong ownership (Newton-Schulz runs on whole - # matrices on one rank; param sync via ``allgather_params``). ``model_chunks`` lets the - # ping-pong fallback break equal-numel ties by a rank-independent identity (see below). - self.shard_params(optimizers, full_param_layouts, model_chunks) - - # Engage FP8 param sync automatically when the decouple-managed params are actually - # quantized (fp8_param_gather on + TE Float8/MXFP8 weights). Off -> plain bf16 path. - # Also tag the gathered fp8 params: the fp8 all-gather (``_allgather_helper_fp8``) - # requantizes bf16 -> each rank's fp8 ``param.data``, so the child optimizer's pre-gather - # fp8 copy-back into ``param.data`` is redundant for them and is skipped. Params in these - # per-rank lists are all-gathered (dp_cp / expt_dp size > 1 here); non-gathered fp8 params - # (e.g. expt_dp == 1 experts, which are absent from these lists) still need the copy-back. - self.use_fp8_param_sync = False - if self.decouple_ddp_layout: - for params_list in (self.dp_cp_params_list, self.expt_dp_params_list): - if not params_list: - continue - for per_rank in params_list: - for p in per_rank: - if is_float8tensor(p): - self.use_fp8_param_sync = True - p._layer_wise_fp8_gathered = True - - # When a full_param_layout is available (and we are not decoupling), - # ddp_config.use_distributed_optimizer is True and model params are views into the - # DDP param buffer. After the optimizer step copies updated fp32 main params → bf16 - # model params, the buffer is already up-to-date in-place. We can use DDP's - # buffer-based all-gather (start_param_sync) instead of the flatten/unflatten - # allgather_params path. In the decouple path (incl. FP8 param-gather), Muon buffers are - # non-DistOpt and own whole params via ping-pong, so we use the legacy allgather_params. + self.shard_params(optimizers, full_param_layouts) + + # When a full_param_layout is available, ddp_config.use_distributed_optimizer + # is True and model params are views into the DDP param buffer. After the + # optimizer step copies updated fp32 main params → bf16 model params, the + # buffer is already up-to-date in-place. We can use DDP's buffer-based + # all-gather (start_param_sync) instead of the flatten/unflatten allgather_params + # path. self.use_buffer_param_sync = full_param_layouts is not None # Set up overlap param gather using DDP bucket infrastructure. @@ -474,27 +513,13 @@ def __init__( optimizers[i] = Float16OptimizerWithFloat16Params( opt, config, None, init_state_fn_list[i] if init_state_fn_list else None ) - # Non-DistOpt LayerWise child has no byte-shard param buffer, so tag it to route - # step_with_ready_grads to copy fp32 master straight into model ``param.data`` even - # under reuse_grad_buf (which would otherwise call the unsupported - # ``_copy_main_params_to_param_buffer``); the gather staging reuses the grad buffer. - optimizers[i]._layer_wise_non_distopt_child = True - - # shard_params() removed non-owned params from the local optimizer groups, so the - # Float16 wrapping above only clears the TE high-precision init copy (a full-size CPU - # tensor per fp8 param) for locally owned params. Without this sweep every DP rank - # retains ~(dp-1)/dp of the LayerWise matrix params' bf16 CPU copies for the whole - # run. The per-rank ownership lists cover all gathered params (owned entries were - # already cleared during master creation; TE's clear is a no-op then). Scoped to - # LayerWise-managed params only: sibling DistOpt params must keep their init val - # until their own optimizer's master creation consumes it. - for params_list in (self.dp_cp_params_list, self.expt_dp_params_list): - if not params_list: - continue - for per_rank_params in params_list: - for p in per_rank_params: - if hasattr(p, 'clear_high_precision_init_val'): - p.clear_high_precision_init_val() + + self.tp_group = self.pg_collection.tp + self.expert_tp_group = getattr(self.pg_collection, 'expt_tp', self.tp_group) + for optimizer in optimizers: + # Child optimizers perform TP duplicate filtering when collecting gradients. + optimizer.tp_group = self.tp_group + optimizer.expert_tp_group = self.expert_tp_group super().__init__(optimizers) @@ -512,7 +537,7 @@ def __init__( # This way each rank do some duplicated work but allgather_v is no longer needed # All current distopt optimization can also be potentially applied - def shard_params(self, optimizers, full_param_layouts=None, model_chunks=None): + def shard_params(self, optimizers, full_param_layouts=None): """Shard params across ranks according to the computed param layout. Each param's shard assignment is derived from the :class:`FullParamLayout` @@ -530,23 +555,23 @@ def shard_params(self, optimizers, full_param_layouts=None, model_chunks=None): chunk). ``None`` triggers the legacy fallback. """ # Simplify when dp_cp group size is 1. - dp_cp_size = get_pg_size(self.pg_collection.dp_cp) + dp_cp_size = get_pg_size(self.dp_cp) if dp_cp_size == 1: self.dp_cp_params_list = None self.expt_dp_params_list = None return - expt_dp_size = get_pg_size(self.pg_collection.expt_dp) + expt_dp_size = get_pg_size(self.expt_dp) if full_param_layouts is not None: self._shard_params_from_layout(optimizers, full_param_layouts, dp_cp_size, expt_dp_size) else: - self._shard_params_ping_pong(optimizers, dp_cp_size, expt_dp_size, model_chunks) + self._shard_params_ping_pong(optimizers, dp_cp_size, expt_dp_size) def _shard_params_from_layout(self, optimizers, full_param_layouts, dp_cp_size, expt_dp_size): """Derive shard assignments from the param layout.""" - dp_cp_rank = get_pg_rank(self.pg_collection.dp_cp) - expt_dp_rank = get_pg_rank(self.pg_collection.expt_dp) + dp_cp_rank = get_pg_rank(self.dp_cp) + expt_dp_rank = get_pg_rank(self.expt_dp) self.dp_cp_params_list = [[] for _ in range(dp_cp_size)] self.expt_dp_params_list = [[] for _ in range(expt_dp_size)] @@ -560,14 +585,15 @@ def _shard_params_from_layout(self, optimizers, full_param_layouts, dp_cp_size, # separate DistributedOptimizer; LayerWise does not own them. if not buffer_key.is_managed_by_layer_wise_optimizer: continue - dp_size = expt_dp_size if buffer_key.is_expert_parallel else dp_cp_size for param, ( param_start_index, param_end_index, bucket_id, ) in layout.param_index_map.items(): bucket_start_index, bucket_end_index = layout.bucket_indices[bucket_id] - shard_size = (bucket_end_index - bucket_start_index) // dp_size + shard_size = ( + bucket_end_index - bucket_start_index + ) // layout.num_optimizer_shards shard_id = (param_start_index - bucket_start_index) // shard_size shard_end_index = bucket_start_index + (shard_id + 1) * shard_size assert param_end_index <= shard_end_index, ( @@ -610,44 +636,16 @@ def _shard_params_from_layout(self, optimizers, full_param_layouts, dp_cp_size, if expt_dp_size == 1 or len(self.expt_dp_params_list[0]) == 0: self.expt_dp_params_list = None - def _build_param_sort_keys(self, model_chunks): - """Build ``{param: (chunk_idx, buffer_idx, global_start_index)}`` — a rank-independent key - for every requires-grad param. - - Both the chunk/buffer enumeration order and the ``param_index_map`` offsets come purely - from model construction (identical across DP ranks), so the key is the same on every rank. - Used to break equal-numel ties in ``_shard_params_ping_pong``. Returns ``None`` if no layout - info is available, so the caller falls back to legacy numel-only ordering. - """ - if model_chunks is None: - return None - identity: Dict[torch.nn.Parameter, tuple] = {} - for chunk_idx, chunk in enumerate(model_chunks): - buffers = list(getattr(chunk, 'buffers', [])) + list( - getattr(chunk, 'expert_parallel_buffers', []) - ) - for buffer_idx, buffer in enumerate(buffers): - param_index_map = getattr(buffer, 'param_index_map', None) - if param_index_map is None: - continue - for param, (global_start, _global_end, _bucket_id) in param_index_map.items(): - identity[param] = (chunk_idx, buffer_idx, global_start) - return identity or None - - def _shard_params_ping_pong(self, optimizers, dp_cp_size, expt_dp_size, model_chunks=None): - """Legacy ping-pong shard assignment (no layout available). + def _shard_params_ping_pong(self, optimizers, dp_cp_size, expt_dp_size): + """Legacy ping-pong-by-numel shard assignment (no layout available). Legacy: this method is a fallback for when no ``full_param_layout`` is provided. Once all call sites supply a layout, this can be removed in favor of :meth:`_shard_params_from_layout`. - Parameters are sorted by a rank-independent TOTAL order and assigned ping-pong style. E.g. - 4 ranks, 10 params p0-p9 -> [[p0, p7, p8], [p1, p6, p9], [p2, p5], [p3, p4]]. - - CRITICAL: the sort key MUST be identical across DP ranks. ``numel`` alone is not (stable - sort tie-breaks equal-numel params by insertion order), which would give different owners - per rank -> params double-owned or zero-owned on the first step. So we tie-break by the - canonical identity ``(chunk_idx, buffer_idx, global_start_index)``. + List of parameters are sorted by numel and assigned to ranks in ping-pong style. + Example of 4 ranks and 10 parameters p0-p9 after sorting, then dp_cp_params_list + will be [[p0, p7, p8], [p1, p6, p9], [p2, p5], [p3, p4]]. """ dp_cp_idx, expt_dp_idx = 0, 0 # Create ping-pong style loop so memory is more balanced. @@ -660,36 +658,23 @@ def _shard_params_ping_pong(self, optimizers, dp_cp_size, expt_dp_size, model_ch for optimizer in optimizers: param_groups += optimizer.param_groups - # Sort param in all groups by a rank-independent TOTAL order, then assign to each rank. - identity = self._build_param_sort_keys(model_chunks) + # Sort param in all groups by param numel and assign to each rank evenly. param_list = [] for group_index, group in enumerate(param_groups): for p in group["params"]: param_list.append((p, group_index)) - if identity is not None: - # Total order: (numel, canonical-global-identity). Identical on every DP rank. - missing = [p for (p, _) in param_list if p not in identity] - assert not missing, ( - "ping-pong ownership requires a canonical identity for every Muon param, " - f"but {len(missing)} param(s) were not found in any model-chunk buffer's " - "param_index_map. Cannot guarantee identical ownership across ranks (the " - "allgather_params gather assumes every rank agrees on each param's single owner)." - ) - param_list.sort(key=lambda x: _param_sort_key(x[0].numel(), identity[x[0]])) - else: - # No layout info: keep the legacy numel-only ordering. - param_list.sort(key=lambda x: x[0].numel()) + param_list.sort(key=lambda x: x[0].numel()) param_groups_this_rank = [[] for g in param_groups] # Assign params to rank in ping-pong style loop. for p, group_index in param_list: if param_groups[group_index].get("is_expert_parallel", False): - if expt_dp_loop[expt_dp_idx] == get_pg_rank(self.pg_collection.expt_dp): + if expt_dp_loop[expt_dp_idx] == get_pg_rank(self.expt_dp): param_groups_this_rank[group_index].append(p) self.expt_dp_params_list[expt_dp_loop[expt_dp_idx]].append(p) expt_dp_idx = (expt_dp_idx + 1) % len(expt_dp_loop) else: - if dp_cp_loop[dp_cp_idx] == get_pg_rank(self.pg_collection.dp_cp): + if dp_cp_loop[dp_cp_idx] == get_pg_rank(self.dp_cp): param_groups_this_rank[group_index].append(p) self.dp_cp_params_list[dp_cp_loop[dp_cp_idx]].append(p) dp_cp_idx = (dp_cp_idx + 1) % len(dp_cp_loop) @@ -723,9 +708,7 @@ def set_bucket_layerwise_params_list(self, model_chunks): if not _bucket_is_managed_by_layer_wise_optimizer(bucket): continue if self.dp_cp_params_list is not None: - bucket_params_list = [ - [] for _ in range(get_pg_size(self.pg_collection.dp_cp)) - ] + bucket_params_list = [[] for _ in range(get_pg_size(self.dp_cp))] for bucket_list, full_params_list in zip( bucket_params_list, self.dp_cp_params_list ): @@ -733,8 +716,8 @@ def set_bucket_layerwise_params_list(self, model_chunks): if param in bucket.params: bucket_list.append(param) else: - # dp_cp_size == 1: single rank owns all params; init the structure anyway - # (mirrors the expert block; shard_params sets dp_cp_params_list=None here). + # dp_cp_size == 1: single rank owns all params, no + # all-gather needed but data structures must be initialized. bucket_params_list = [list(bucket.params_list)] bucket.set_layerwise_params_list(bucket_params_list) # Do the same for expert parallel bucket groups. @@ -743,9 +726,7 @@ def set_bucket_layerwise_params_list(self, model_chunks): if not _bucket_is_managed_by_layer_wise_optimizer(bucket): continue if self.expt_dp_params_list is not None: - bucket_params_list = [ - [] for _ in range(get_pg_size(self.pg_collection.expt_dp)) - ] + bucket_params_list = [[] for _ in range(get_pg_size(self.expt_dp))] for bucket_list, full_params_list in zip( bucket_params_list, self.expt_dp_params_list ): @@ -766,70 +747,8 @@ def allgather_params(self) -> None: call sites supply a ``full_param_layout``, this can be removed — the standard distributed optimizer buffer all-gather (via ``start_param_sync``) replaces this flatten/unflatten path. - - Two transport variants share the same uneven (all-gather-v) shape: - - * **bf16** (``use_fp8_param_sync=False``): all-gather owned bf16 ``param.data``, copy_ into - non-owned params. - * **fp8** (``use_fp8_param_sync=True``): stage owned fp32 master->bf16, all-gather bf16, - requantize into EVERY rank's ``param.data`` (owned included) so all hold - ``Q(bf16(master))`` (== OFF/Adam). Then ``post_all_gather_processing`` rebuilds fp8 - columnwise/transpose (blockwise/Float8; mxfp8 noop since copy-back already forced it). """ - # FP8-aware variant: stage bf16, uneven all-gather bf16, requantize per rank. - def _allgather_helper_fp8(params_list, group): - # TODO(perf, blockwise-only): blockwise could gather the owner's fp8 rowwise data - # (~2x less comm) instead of bf16; mxfp8 must stay on bf16. See the matching TODO in - # ``_ParamAndGradBucketGroup.start_param_sync`` for the full rationale. - rank = get_pg_rank(group) - dp_size = get_pg_size(group) - # Device from any non-empty owned list (rank 0 may own zero params in the layout). - device = next((params[0].device for params in params_list if len(params) > 0), None) - if device is None: - # No rank owns any param in this buffer -> nothing to gather. - return - - # Stage fp32 master->bf16 (high-precision source), not lossy dequant(fp8). - owned = params_list[rank] - src = ( - _flatten_dense_tensors([_stage_param_to_bf16(p) for p in owned]) - if len(owned) > 0 - else torch.empty(0, device=device, dtype=torch.bfloat16) - ) - flat_sizes = [sum(p.numel() for p in params) for params in params_list] - if max(flat_sizes) == 0: - return - - gather_list = [] - for i in range(dp_size): - if i == rank: - gather_list.append(src) - else: - gather_list.append( - torch.empty(flat_sizes[i], device=device, dtype=torch.bfloat16) - ) - - torch.distributed.all_gather(gather_list, src, group=group) - - # Requantize the gathered bf16 into EVERY rank's params (owned included) so all ranks - # hold Q(bf16(master)), matching OFF/Adam. Unflatten by param shape (logical numel). - for idx, params in enumerate(params_list): - if len(params) == 0: - continue - templates = [ - torch.empty(p.shape, device="meta", dtype=torch.bfloat16) for p in params - ] - updated_params = _unflatten_dense_tensors(gather_list[idx], templates) - for updated_bf16, model_p in zip(updated_params, params): - copy_back_gathered_bf16_into_fp8_param(model_p, updated_bf16) - - # Rebuild fp8 columnwise/transpose after the gather (mirrors the overlap / DistOpt - # paths; blockwise/Float8 build it, mxfp8 is a noop). Else it'd be deferred to forward. - fp8_params = [p for params in params_list for p in params if is_float8tensor(p)] - if fp8_params: - post_all_gather_processing(fp8_params) - # helper function to flatten local params, all-gather, # unflatten and copy to model params def _allgather_helper(params_list, group): @@ -869,35 +788,10 @@ def _allgather_helper(params_list, group): if self.pg_collection is None: return - - def _dispatch(params_list, group): - # Split each rank's owned params by transport dtype. fp8 and bf16 params ride the - # bf16 transport (the fp8 helper stages master->bf16 / requantizes, a no-op in - # precision for bf16); native fp32 params (e.g. weights marked keep_in_fp32 such as - # the DeepSeek-V4 CSA ``ape``) must be gathered in fp32 -- routing them through the - # bf16-staged path would silently downcast them, and mixing fp32 with bf16 in one - # flatten is invalid. For a pure-bf16 model (no fp32 Muon params) the native group is - # empty and this collapses to the original single-helper dispatch. - staged = [ - [p for p in owned if is_float8tensor(p) or p.dtype != torch.float32] - for owned in params_list - ] - native = [ - [p for p in owned if not is_float8tensor(p) and p.dtype == torch.float32] - for owned in params_list - ] - if any(owned for owned in staged): - staged_helper = ( - _allgather_helper_fp8 if self.use_fp8_param_sync else _allgather_helper - ) - staged_helper(staged, group) - if any(owned for owned in native): - _allgather_helper(native, group) - if self.dp_cp_params_list: - _dispatch(self.dp_cp_params_list, self.pg_collection.dp_cp) + _allgather_helper(self.dp_cp_params_list, self.dp_cp) if self.expt_dp_params_list: - _dispatch(self.expt_dp_params_list, self.pg_collection.expt_dp) + _allgather_helper(self.expt_dp_params_list, self.expt_dp) @torch.no_grad() def broadcast_params(self): @@ -906,15 +800,15 @@ def broadcast_params(self): if self.dp_cp_params_list is None: return for i, params in enumerate(self.dp_cp_params_list): - src_global_rank = torch.distributed.get_global_rank(self.pg_collection.dp_cp, i) + src_global_rank = torch.distributed.get_global_rank(self.dp_cp, i) for p in params: - torch.distributed.broadcast(p, src_global_rank, self.pg_collection.dp_cp) + torch.distributed.broadcast(p, src_global_rank, self.dp_cp) if self.expt_dp_params_list is None: return for i, params in enumerate(self.expt_dp_params_list): - src_global_rank = torch.distributed.get_global_rank(self.pg_collection.expt_dp, i) + src_global_rank = torch.distributed.get_global_rank(self.expt_dp, i) for p in params: - torch.distributed.broadcast(p, src_global_rank, self.pg_collection.expt_dp) + torch.distributed.broadcast(p, src_global_rank, self.expt_dp) @torch.no_grad() def get_grad_norm(self): @@ -970,6 +864,8 @@ def count_zeros(self): params, grad_stats_parallel_group=None, use_decoupled_grad=self.config.use_precision_aware_optimizer_no_fp8_or_ds_fp8, + tp_group=self.tp_group, + expert_tp_group=self.expert_tp_group, ) def start_param_sync_for_bucket_group_subset(self) -> None: @@ -1047,14 +943,20 @@ def sharded_state_dict( model_sharded_state_dict, is_loading, **kwargs ) + # (E)GTP_remat-replicated -> (gtp_rank, gtp_remat_size), consumed by _fold_replica_id. + gtp_fold = _build_gtp_replica_fold(self.pg_collection, self.model_chunks) + # for fixed DP usage only for sh_base in nested_values(sharded_state_dict): if hasattr(sh_base, 'replica_id'): assert ( isinstance(sh_base.replica_id, int) or len(sh_base.replica_id) == 3 ), f'Expected replica_id as int or (PP, TP, DP), got: {sh_base}' - sh_base.replica_id = ( - 0 if isinstance(sh_base.replica_id, int) else (*sh_base.replica_id[:2], 0) + if isinstance(sh_base.replica_id, int): + sh_base.replica_id = 0 + continue + sh_base.replica_id = _fold_replica_id( + sh_base.replica_id, getattr(sh_base, 'key', ''), gtp_fold ) # later code assume list but chained optimizer fallback to non-list if there's only one @@ -1092,3 +994,14 @@ def sharded_state_dict( nonempty_rank_group['params'] = local_params sd['optimizer']['param_groups'][i] = nonempty_rank_group return sharded_state_dict + + def save_state_dict_to_file(self, filename: str) -> None: + """Save the parameter state of the optimizer. For torch format only. + Args: + filename: The filename to save the parameter state. + """ + torch.save(super().state_dict(), filename) + + def load_state_dict_from_file(self, filename: str) -> None: + """Load the parameter state of the optimizer. For torch format only.""" + super().load_state_dict(torch.load(filename)) diff --git a/megatron/core/optimizer/optimizer.py b/megatron/core/optimizer/optimizer.py index b5532018c61..3821290ebdf 100644 --- a/megatron/core/optimizer/optimizer.py +++ b/megatron/core/optimizer/optimizer.py @@ -7,6 +7,7 @@ import math import warnings from abc import ABC, abstractmethod +from itertools import chain from logging import getLogger from typing import Any, Callable, Dict, List, Optional, Tuple, Union @@ -46,6 +47,7 @@ ) from ..dist_checkpointing.utils import add_prefix_for_sharding from ..fp8_utils import copy_back_gathered_bf16_into_fp8_param, is_float8tensor +from ..optimizer_param_scheduler import ParamGroupOverride as _ParamGroupOverride from ..transformer.module import param_is_not_shared from ..utils import log_single_rank from .clip_grads import clip_grad_by_total_norm_fp32, count_zeros_fp32, get_grad_norm_fp32 @@ -94,7 +96,58 @@ def _multi_tensor_copy_this_to_that( that_.copy_(this_) -param_group_identifier_keys = ('wd_mult', 'lr_mult', 'is_expert_parallel', 'is_decoupled_lr') +# Per-group keys used to uniquely identify a param_group during save/load matching. +# Used by ``DistributedOptimizer.load_state_dict`` and +# ``MegatronOptimizer._filter_and_reorder_param_groups`` to map saved param_groups +# onto current param_groups by behavioral equivalence. +# +# This MUST cover every per-group field that influences scheduler or optimizer behavior; +# otherwise two groups that differ only in a missing key (e.g. ``max_lr``) will collide +# in the matching dict and one will silently overwrite the other on load. That's a +# correctness bug: the load returns silently but the LR/WD applied at the next +# optimizer step is wrong, leading to loss explosion on a converged-enough model. +# +# Source of truth for the user-overridable fields is +# :class:`megatron.core.optimizer_param_scheduler.ParamGroupOverride` (the keys the +# scheduler reads from each param_group via ``param_group.get(...)``): +# ``max_lr``, ``min_lr``, ``start_wd``, ``end_wd``, ``wd_mult``, ``optimizer``. +# We pull those keys directly from that TypedDict's annotations so future +# additions to ``ParamGroupOverride`` automatically extend the identifier. +# +# The remaining keys (``lr_mult``, ``is_expert_parallel``, ``is_decoupled_lr``) are +# structural flags set by ``_get_param_groups`` (in this module's ``__init__.py``) +# on every param_group at construction time. They aren't part of ``ParamGroupOverride`` +# (users don't override them directly; they're implied by ``decoupled_lr`` config and +# expert-parallel sharding), so we list them explicitly. +def _param_group_override_keys() -> tuple[str, ...]: + """Return every field declared on ``ParamGroupOverride``. + + For any TypedDict, ``__annotations__.keys() == __required_keys__ | __optional_keys__`` + - the (required, optional) pair is just a *partition* of the declared set + based on the TypedDict's totality choice and any ``Required[]`` / + ``NotRequired[]`` wrappers. We want the whole declared set, regardless of + how it's partitioned, so we read ``__annotations__`` directly. Reading just + one side of the partition would silently miss fields if a future maintainer + flipped totality or introduced wrappers: + + total=False (current): optional={max_lr, min_lr, ...}, required={} + total=True: required={max_lr, min_lr, ...}, optional={} + mixed Required/NotRequired: fields split between the two sides + """ + return tuple(sorted(_ParamGroupOverride.__annotations__.keys())) + + +param_group_identifier_keys = ( + # Per-group user-overridable keys (single source of truth: ParamGroupOverride). + # The scheduler reads ``max_lr``/``min_lr`` in ``get_lr`` and ``start_wd``/``end_wd`` + # in ``get_wd``; ``wd_mult`` is multiplied into ``weight_decay`` in ``step``; + # ``optimizer`` selects per-group optimizer class. + *_param_group_override_keys(), + # Optimizer-side structural flags (not user-overridable via ParamGroupOverride): + 'lr_mult', + 'is_expert_parallel', + 'is_decoupled_lr', +) MTP_GRAD_NORM_GROUP = 'mtp' GRAD_NORM_GROUP_ATTR = 'grad_norm_group' SEPARATE_GRAD_NORM_GROUPS = (MTP_GRAD_NORM_GROUP,) @@ -124,6 +177,8 @@ def _is_separate_grad_norm_group(grad_norm_group: Optional[str]) -> bool: def copy_optimizer_param_metadata(destination: torch.Tensor, source: torch.Tensor) -> None: """Copy optimizer-relevant metadata when creating param views/copies.""" + if hasattr(source, 'allreduce'): + destination.allreduce = source.allreduce if hasattr(source, 'shared'): destination.shared = source.shared if hasattr(source, GRAD_NORM_GROUP_ATTR): @@ -187,6 +242,7 @@ def _filter_grads_for_norm( - parameter should not be shared (i.e., grads shouldn't be double counted while computing norms). - should not be a replica due to tensor model parallelism. + - should not be a replica due to (expert) generalized tensor parallelism. """ grads_for_norm = [] for param in params: @@ -213,9 +269,12 @@ def _filter_grads_for_norm( grad_not_none = grad is not None is_not_shared = param_is_not_shared(param) is_not_tp_duplicate = tensor_parallel.param_is_not_tensor_parallel_duplicate( - param, getattr(self, 'tp_group', None) + param, + tp_group=getattr(self, 'tp_group', None), + expert_tp_group=getattr(self, 'expert_tp_group', None), ) - if grad_not_none and is_not_shared and is_not_tp_duplicate: + is_not_gtp_duplicate = tensor_parallel.param_is_not_gtp_duplicate(param) + if grad_not_none and is_not_shared and is_not_tp_duplicate and is_not_gtp_duplicate: grads_for_norm.append(grad) return grads_for_norm @@ -380,6 +439,7 @@ def count_zeros(self) -> float: and getattr(params[0], "__fsdp_param__", False) ), tp_group=getattr(self, 'tp_group', None), + expert_tp_group=getattr(self, 'expert_tp_group', None), ) @abstractmethod @@ -535,9 +595,10 @@ def restore_from_cpu(self): def _filter_and_reorder_param_groups( current_groups: List[Dict], state_dict_groups: List[Dict] ) -> List[Dict]: - """Filter and reorder state_dict parameter groups to match current optimizer groups. - Keys used for matching align with those from _get_param_groups: - (wd_mult, lr_mult, is_expert_parallel, is_decoupled_lr) + """Pair each current param_group with its saved counterpart by identifier tuple. + + Construction order isn't part of the checkpoint, so we match by a tuple of + per-group config (``param_group_identifier_keys``) rather than by position. Args: current_groups (List[Dict]): Parameter groups from the current optimizer instance. @@ -549,24 +610,29 @@ def _filter_and_reorder_param_groups( Raises: ValueError: If parameter groups in state dict don't match current optimizer. """ - # Define groups order that is needed in the current optimizer (coming from runtime) - needed_groups = [ - # NeMo may have different key for required fields, e.g., "wd_mult" to "pre_wd_mult" - tuple(g[key] if key in g else g[f"pre_{key}"] for key in param_group_identifier_keys) - for g in current_groups - ] - # Keep state_dict param group order since groups are LocalNonpersistentObject - # and their order is determined at runtime, not from the checkpoint. + def _identifier_for(group: dict) -> tuple: + out = [] + for key in param_group_identifier_keys: + # NeMo aliases ``wd_mult``/``lr_mult`` as ``pre_wd_mult``/``pre_lr_mult``. + if key in group: + out.append(group[key]) + elif f"pre_{key}" in group: + out.append(group[f"pre_{key}"]) + else: + # Treat missing and explicit None identifier values as equivalent. + out.append(None) + return tuple(out) + + needed_groups = [_identifier_for(g) for g in current_groups] params_in_state_dict_order = [g['params'] for g in state_dict_groups] - loaded_groups_map = { - tuple( - # NeMo may have different key for required fields, e.g., "wd_mult" to "pre_wd_mult" - group[key] if key in group else group[f"pre_{key}"] - for key in param_group_identifier_keys - ): group - for group in state_dict_groups - } + # Duplicate identifiers here silently clobber: two saved groups with the same tuple + # collapse to whichever was inserted last, and one current group inherits the wrong + # override state (``max_lr`` etc.). Params are unaffected — they come from the + # current optimizer below — but the next step runs at the wrong LR / WD. Adding the + # distinguishing field to ``param_group_identifier_keys`` is the fix. See + # ``test_filter_reorder_distinguishes_groups_by_max_lr``. + loaded_groups_map = {_identifier_for(group): group for group in state_dict_groups} final_groups = [] for key, params in zip(needed_groups, params_in_state_dict_order): @@ -731,9 +797,14 @@ def step_with_ready_grads(self) -> bool: barrier=self.config.barrier_with_L1_time ) if not self.is_stub_optimizer: + # The reuse_grad_buf (fp8-param-gather) path stages master params into the DDP + # param buffer, which only DistributedOptimizer owns. Optimizers without it + # (e.g. LayerWiseDistributedOptimizer's Float16 base opts) must instead copy + # master -> model params so the forward sees the update. if ( self.config.reuse_grad_buf_for_mxfp8_param_ag and not self._layer_wise_non_distopt_child + and hasattr(self, "_copy_main_params_to_param_buffer") ): # In the case of overlap_param_gather, # copy is manually called in the training loop @@ -786,6 +857,120 @@ def step(self): return success, grad_norm, num_zeros_in_grad +def _strip_module_prefix(name: str) -> str: + """Strip wrapper ``module.`` prefixes (DDP/Float16Module) off a dotted param name.""" + while name.startswith('module.'): + name = name[len('module.') :] + return name + + +def _backfill_gtp_sharded_param_map( + id_to_sharded_param_map: dict, float16_groups, model_sharded_state_dict=None +) -> None: + """Backfill the optimizer id->ShardedTensor map with GTP_remat shards it is missing (in place). + + WHAT: ``get_param_id_to_sharded_param_map`` matches an optimizer param to its model + ShardedTensor by object identity (``id(model_entry.data) == id(optim_param)``). Two GTP_remat + cases break that match: + 1. Native-FP8 GTP weights: the model entry's data is a *dequantized BF16 copy* of the param + (make_tp_sharded_tensor_for_checkpoint). The copy carries a ``_gtp_dequant_src`` backlink + to the live FP8 param, so the model's OWN entry is reused here (identity first, tagged + ``_debug_name`` second) -- preserving its full offsets (expert axes included) and + replica_id. + 2. Gathered+split factory params (Mamba ``in_proj``): the model entry exposes the *gathered* + tensor, so nothing matches the per-shard GTP param. Rebuild the same per-shard + ShardedTensor every other GTP_remat weight gets. The rebuild is NOT expert-parallel + aware (no expert offsets/replica), so expert params must resolve via case 1; refuse + loudly instead of writing colliding shards across EP groups. + + WHEN: only the distributed-Muon path reaches here. ``LayerWiseDistributedOptimizer`` keeps such + matrix params whole and routes them through this ``Float16OptimizerWithFloat16Params``. + Distributed Adam uses its own ``DistributedOptimizer.sharded_state_dict`` (flat-buffer path) + and is unaffected. + + No-op when GTP is unavailable or when every param already matched. + """ + try: + from megatron.core.tensor_parallel.gtp_api import ( + is_gtp_param, + make_sharded_tensors_for_checkpoint_with_gtp_remat, + ) + except ImportError: + return # GTP not built in -- nothing to backfill. + + # is_gtp_param matches both the legacy BF16 slice params and native-FP8 GTP params. + unmatched = [ + (param_id, p) + for param_id, p in enumerate(chain.from_iterable(float16_groups)) + if param_id not in id_to_sharded_param_map and is_gtp_param(p) + ] + if not unmatched: + return + + from ..dist_checkpointing.dict_utils import nested_values + from ..dist_checkpointing.mapping import ShardedTensor + + # Index the model's own entries by (a) the dequantized-copy backlink and (b) checkpoint key. + src_id_to_entry = {} + key_to_entry = {} + if model_sharded_state_dict is not None: + for entry in nested_values(model_sharded_state_dict): + src = getattr(getattr(entry, 'data', None), '_gtp_dequant_src', None) + if src is not None: + src_id_to_entry[id(src)] = entry + key = getattr(entry, 'key', None) + if key is not None: + # Grouped-expert entries share one key (offsets differ) -> ambiguous, drop. + key_to_entry[key] = None if key in key_to_entry else entry + + # Groups sourced lazily (below) only when a rebuild is needed, so GTP models on + # explicit grids (e.g. MiMo) don't require the global MPU groups unless they hit it. + tp_group = None + dp_cp_gtp_remat_group = None + for param_id, p in unmatched: + # Case 1: reuse the model's own entry (native-FP8 dequantized copy broke the id match). + entry = src_id_to_entry.get(id(p)) + if entry is None: + name = _strip_module_prefix(getattr(p, '_debug_name', '') or '') + candidate = key_to_entry.get(name) + # Reuse only a plain ShardedTensor with this shard's local shape; a factory + # (gathered data, e.g. Mamba in_proj) must take the per-shard rebuild below. + if ( + candidate is not None + and isinstance(candidate, ShardedTensor) + and tuple(candidate.data.shape) == tuple(p.shape) + ): + entry = candidate + if entry is not None: + id_to_sharded_param_map[param_id] = entry + continue + # Case 2: rebuild. Not EP-aware -- an expert param rebuilt here would collide across + # expert-parallel groups (duplicate writers), so it must have matched above. + if not getattr(p, 'allreduce', True): + raise ValueError( + f"GTP expert-parallel param '{getattr(p, '_debug_name', '')}' (id {param_id}) " + "has no matching model ShardedTensor; refusing the EP-unaware rebuild (it would " + "write duplicate shards across expert-parallel groups)." + ) + if tp_group is None: + tp_group = parallel_state.get_tensor_model_parallel_group() + # Required kwarg, unused for GTP-sharded params (offset/replica from the gtp axis). + dp_cp_gtp_remat_group = parallel_state.get_data_parallel_group( + with_context_parallel=True + ) + # Key by the param's dotted name (set in prod by tag_gtp_params_with_names); the fallback + # keeps the function usable in tests where the name was not tagged. + key = p._debug_name or f'_gtp_optim_param_{param_id}' + rebuilt = make_sharded_tensors_for_checkpoint_with_gtp_remat( + {key: p}, + prefix='', + tensor_parallel_layers_axis_map={key: 0}, + tp_group=tp_group, + dp_cp_group=dp_cp_gtp_remat_group, + ) + id_to_sharded_param_map[param_id] = rebuilt[key] + + class Float16OptimizerWithFloat16Params(MixedPrecisionOptimizer): """Float16 optimizer for fp16 and bf16 data types. @@ -849,6 +1034,7 @@ def __init__( main_param = param.detach().clone().float() # Copy tensor model parallel attributes. tensor_parallel.copy_tensor_model_parallel_attributes(main_param, param) + tensor_parallel.copy_gtp_attributes(main_param, param) copy_optimizer_param_metadata(main_param, param) # Replace the optimizer params with the new fp32 copy. param_group['params'][i] = main_param @@ -1037,6 +1223,10 @@ def model_params_in_optimizer_order(): model_sharded_state_dict, model_params_in_optimizer_order() ) + _backfill_gtp_sharded_param_map( + id_to_sharded_param_map, self.float16_groups, model_sharded_state_dict + ) + # Convert fp32_from_fp16_params assert len(state_dict['fp32_from_fp16_params']) == len( state_dict['optimizer']['param_groups'] @@ -1695,6 +1885,8 @@ def count_zeros(self): self.config.use_precision_aware_optimizer and getattr(params[0], "__fsdp_param__", False) ), + tp_group=getattr(self.chained_optimizers[0], 'tp_group', None), + expert_tp_group=getattr(self.chained_optimizers[0], 'expert_tp_group', None), ) else: num_zeros_in_grad = 0 @@ -1823,7 +2015,12 @@ def step(self): use_decoupled_grad=use_decoupled_grad, ) - if grad_norm > optimizer.config.grad_norm_skip_threshold and main_params: + grad_norm_skip_threshold = optimizer.config.grad_norm_skip_threshold + if ( + main_params + and math.isfinite(grad_norm_skip_threshold) + and grad_norm > grad_norm_skip_threshold + ): log_single_rank( logger, logging.INFO, "skipping grad norm because it's too large %s", grad_norm ) diff --git a/megatron/core/optimizer/param_layout.py b/megatron/core/optimizer/param_layout.py index 2ee511c6126..9d2dd4db365 100644 --- a/megatron/core/optimizer/param_layout.py +++ b/megatron/core/optimizer/param_layout.py @@ -11,7 +11,7 @@ import math from dataclasses import dataclass, field -from typing import Dict, List, Tuple +from typing import Dict, List, Optional, Tuple import torch @@ -79,12 +79,16 @@ class PerBufferParamLayout: param_indices: The index of each param among same-dtype params (using the "fake" high-precision dtype for FP8/NVFP4 params). Needed for loading non-native-fp8 checkpoints in native-fp8 mode. Order matches param_index_map iteration order. + num_optimizer_shards: Number of optimizer shards. Set by the distributed optimizer + that computes the layout so that shard assignment at runtime uses the same + value. ``None`` for non-distributed-optimizer layouts. """ param_index_map: Dict[torch.nn.Parameter, Tuple[int, int, int]] = field(default_factory=dict) bucket_indices: List[Tuple[int, int]] = field(default_factory=list) per_bucket_numel_unpadded: List[int] = field(default_factory=list) param_indices: List[int] = field(default_factory=list) + num_optimizer_shards: Optional[int] = None @dataclass diff --git a/megatron/core/package_info.py b/megatron/core/package_info.py index 3881c5d4052..7ba1aa93c5f 100644 --- a/megatron/core/package_info.py +++ b/megatron/core/package_info.py @@ -3,7 +3,7 @@ MAJOR = 0 -MINOR = 19 +MINOR = 20 PATCH = 0 PRE_RELEASE = '' diff --git a/megatron/core/parallel_state.py b/megatron/core/parallel_state.py index 70234884ba1..256ff2568f9 100644 --- a/megatron/core/parallel_state.py +++ b/megatron/core/parallel_state.py @@ -27,6 +27,9 @@ # Intra-layer model parallel group that the current rank belongs to. _TENSOR_MODEL_PARALLEL_GROUP = None +# Generalized tensor parallelism group that the current rank belongs to. +_GTP_WEIGHT_REMAT_GROUP = None +_GTP_WEIGHT_REMAT_GLOBAL_RANKS = None # Inter-layer model parallel group that the current rank belongs to. _PIPELINE_MODEL_PARALLEL_GROUP = None # Model parallel group (both intra- and pipeline) that the current rank belongs to. @@ -50,6 +53,9 @@ # _EXPERT_TENSOR denotes tensor parallelism of expert which splits tensor across the group. # _EXPERT_DATA denotes data parallelism of expert which replicates weight across the group. +# Expert generalized tensor parallelism group that current rank belongs to. +_EXPERT_GTP_WEIGHT_REMAT_GROUP = None +_EXPERT_GTP_WEIGHT_REMAT_GLOBAL_RANKS = None # Expert model parallel group that current rank belongs to. _EXPERT_MODEL_PARALLEL_GROUP = None # Expert tensor parallel group that current rank belongs to. @@ -58,12 +64,18 @@ _EXPERT_TENSOR_AND_MODEL_PARALLEL_GROUP = None # Expert tensor, model, pipeline combined parallel group _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP = None +# Same as above, but additionally merged across EGTP peers (analog of dense _MODEL_PARALLEL_GROUP +# under GTP_remat). Identical to the above when EGTP_remat_size=1. +_EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP_WITH_EGTP = None # Expert data parallel group _EXPERT_DATA_PARALLEL_GROUP = None _EXPERT_DATA_PARALLEL_GROUP_GLOO = None _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP = None _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP_GLOO = None _INTER_PARTIAL_EXPERT_DATA_PARALLEL_GROUP = None +# Full expert data-parallel groups: span the egtp_remat axis, for data distribution. +_EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT = None +_INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT = None # Parallel state values changed on the fly _MPU_EXPERT_MODEL_PARALLEL_WORLD_SIZE = None _MPU_EXPERT_MODEL_PARALLEL_RANK = None @@ -118,6 +130,13 @@ # Dynamic context parallel groups _DYNAMIC_DP_CP_GROUPS = {} +# Full data-parallel groups: span every distinct-data rank +# (size = replicate_DP x gtp_remat). Used for data distribution (batch split, num-microbatches, +# gradient scaling) and reductions covering all distinct-data ranks. +_DATA_PARALLEL_GROUP_WITH_GTP_REMAT = None +_DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT = None +_INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT = None + # Data parallel group information with context parallel combined. _DATA_PARALLEL_GROUP_WITH_CP = None _DATA_PARALLEL_GROUP_WITH_CP_GLOO = None @@ -130,7 +149,7 @@ # combined parallel group of TP and CP _TENSOR_AND_CONTEXT_PARALLEL_GROUP = None -# combined parallel group of TP, DP, and CP used for fp8 +# combined parallel group of TP, DP, and CP used for fp8 (spans gtp_remat, like dp) _TENSOR_AND_DATA_PARALLEL_GROUP_WITH_CP = None # Paralel group of all GPUs in a distributed optimizer instance @@ -445,7 +464,15 @@ class RankGenerator(object): """A class for generating rank groups for different modes of parallelism.""" def __init__( - self, tp: int, ep: int, dp: int, pp: int, cp: int, order: str, rank_offset: int = 0 + self, + tp: int, + ep: int, + dp: int, + pp: int, + cp: int, + order: str, + rank_offset: int = 0, + gtp_remat: int = 1, ) -> None: assert ( ep == 1 or cp == 1 @@ -457,8 +484,9 @@ def __init__( self.dp = dp self.pp = pp self.cp = cp + self.gtp_remat = gtp_remat self.rank_offset = rank_offset - self.world_size = tp * dp * pp * cp * ep + self.world_size = tp * dp * pp * cp * ep * gtp_remat self.name_to_size = { "tp": self.tp, @@ -466,6 +494,7 @@ def __init__( "dp": self.dp, "ep": self.ep, "cp": self.cp, + "gtp_remat": self.gtp_remat, } self.order = order order = order.lower() @@ -518,6 +547,13 @@ def get_ranks(self, token): rank_group[i] += self.rank_offset return ranks + def get_gtp_ranks(self, gtp_remat_size: int): + """Get the GTP weight-sharding groups (singletons when ``gtp_remat_size == 1``).""" + assert ( + self.gtp_remat == gtp_remat_size + ), f"gtp_remat axis size ({self.gtp_remat}) != requested gtp_remat_size ({gtp_remat_size})" + return self.get_ranks('gtp_remat') + def default_embedding_ranks(pp_ranks): """Return the default ranks that constitute the stages on which the word embeddings live. @@ -541,6 +577,24 @@ def overwrite_nccl_comm_cfgs(nccl_comm_cfgs, pg_name, key_value_pair): nccl_comm_cfgs[pg_name][key_value_pair[0]] = key_value_pair[1] +def _inject_gtp_remat_axis(order_str: str, after: str = "tp") -> str: + """Inject the 'gtp_remat' axis into a RankGenerator order string for NCCL locality. + + Position controls locality (leftmost token = smallest stride = most adjacent ranks): + - dense/decoder: inject after 'tp' -> 'tp-gtp_remat-cp-ep-dp-pp' (GTP_remat local). + - expert: inject after 'ep' -> 'tp-cp-ep-gtp_remat-dp-pp' so EP keeps more-local placement + than EGTP (the MoE EP all-to-all is the heavier expert-side collective). + When gtp_remat/egtp_remat size is 1 the injected axis is a no-op (singleton groups). + """ + toks = order_str.split("-") + if "gtp_remat" in toks: + return order_str + anchor = after if after in toks else "tp" + pos = (toks.index(anchor) + 1) if anchor in toks else 0 + toks.insert(pos, "gtp_remat") + return "-".join(toks) + + # pylint: disable=C0301 def initialize_model_parallel( tensor_model_parallel_size: int = 1, @@ -553,6 +607,8 @@ def initialize_model_parallel( dynamic_context_parallel: bool = False, min_dynamic_context_parallel_size: int = 1, expert_model_parallel_size: int = 1, + gtp_remat_size: int = 1, + expert_gtp_remat_size: int = 1, num_distributed_optimizer_instances: int = 1, expert_tensor_parallel_size: Optional[int] = None, nccl_communicator_config_path: Optional[str] = None, @@ -632,6 +688,22 @@ def initialize_model_parallel( The number of Mixture of Experts parallel GPUs in each expert parallel group. + gtp_remat_size (int, default = 1): + Generalized tensor parallelism with weight rematerialization (GTP). + Shards model weights along ``out_features`` across this many ranks; + each weight is rematerialized independently (per-weight, not per- + layer) via async all-gather on every forward AND backward pass. A + first-class orthogonal axis (world_size = TP*GTP*CP*DP). Maps to the + dataclass field ``ModelParallelConfig.gtp_weight_remat_size``. + NOTE: "remat" here is NOT activation recomputation/checkpointing. + + expert_gtp_remat_size (int, default = 1): + Expert-side counterpart of ``gtp_remat_size`` — shards routed-expert + weights along ``out_features`` and rematerializes per-weight on + every forward AND backward pass. A first-class orthogonal axis on the + expert grid. Independent from ``gtp_remat_size``. Maps to + ``ModelParallelConfig.expert_gtp_weight_remat_size``. + num_distributed_optimizer_instances (int, default = 1): The number of distributed optimizer replicas across the data- parallel domain. @@ -729,7 +801,22 @@ def initialize_model_parallel( local_world_size if local_world_size is not None else torch.distributed.get_world_size() ) - model_size = tensor_model_parallel_size * pipeline_model_parallel_size * context_parallel_size + # GTP_remat requires a single distributed-optimizer instance: partial-distopt sharding of the + # data domain would need gtp_remat-aware sizing. Assert early so all group builds below can + # assume one instance when GTP_remat/EGTP is active. + assert not ( + (gtp_remat_size > 1 or expert_gtp_remat_size > 1) + and num_distributed_optimizer_instances > 1 + ), "GTP_remat with num_distributed_optimizer_instances > 1 is not yet supported." + + # gtp_remat counts toward model_size (it consumes its own ranks and carries distinct data), + # so data_parallel_size becomes the replicate degree. + model_size = ( + tensor_model_parallel_size + * pipeline_model_parallel_size + * context_parallel_size + * gtp_remat_size + ) if world_size % model_size != 0: raise RuntimeError(f"world_size ({world_size}) is not divisible by {model_size}") @@ -766,21 +853,28 @@ def initialize_model_parallel( for pg_name in high_priority_stream_groups: overwrite_nccl_comm_cfgs(nccl_comm_cfgs, pg_name, ("is_high_priority_stream", True)) + decoder_order = _inject_gtp_remat_axis(order, after="tp") + decoder_rank_generator = RankGenerator( tp=tensor_model_parallel_size, ep=1, dp=data_parallel_size, pp=pipeline_model_parallel_size, cp=context_parallel_size, - order=order, + order=decoder_order, rank_offset=rank_offset, + gtp_remat=gtp_remat_size, ) # Build expert rank generator if expert_tensor_parallel_size is None: expert_tensor_parallel_size = tensor_model_parallel_size + # EGTP is a world-size factor for the expert grid too (mirrors gtp_remat on the dense grid). expert_tensor_model_pipeline_parallel_size = ( - expert_tensor_parallel_size * expert_model_parallel_size * pipeline_model_parallel_size + expert_tensor_parallel_size + * expert_model_parallel_size + * pipeline_model_parallel_size + * expert_gtp_remat_size ) expert_data_parallel_size = world_size // expert_tensor_model_pipeline_parallel_size if world_size % expert_tensor_model_pipeline_parallel_size != 0: @@ -788,15 +882,16 @@ def initialize_model_parallel( f"world_size ({world_size}) is not divisible by expert_tensor_model_pipeline_parallel size ({expert_tensor_model_pipeline_parallel_size})" ) - # TODO: support expert specific ordering + expert_order = _inject_gtp_remat_axis(order, after="ep") expert_decoder_rank_generator = RankGenerator( tp=expert_tensor_parallel_size, ep=expert_model_parallel_size, dp=expert_data_parallel_size, pp=pipeline_model_parallel_size, cp=1, - order=order, + order=expert_order, rank_offset=rank_offset, + gtp_remat=expert_gtp_remat_size, ) assert ( @@ -832,6 +927,29 @@ def initialize_model_parallel( data_parallel_size * context_parallel_size ) // num_distributed_optimizer_instances + # Build the generalized tensor parallel groups. + # GTP_remat overlaps with the CP-DP domain because GTP_remat only shards weights + # while CP only shards activations — they are independent and can share ranks. + global _GTP_WEIGHT_REMAT_GROUP + global _GTP_WEIGHT_REMAT_GLOBAL_RANKS + assert ( + _GTP_WEIGHT_REMAT_GROUP is None + ), "generalized tensor parallel group is already initialized" + for gtp_ranks in decoder_rank_generator.get_gtp_ranks(gtp_remat_size): + group = create_group( + gtp_ranks, + timeout=timeout, + pg_options=get_nccl_options("gtp_remat", nccl_comm_cfgs), + group_desc="GTP_WEIGHT_REMAT_GROUP", + ) + if rank in gtp_ranks: + _GTP_WEIGHT_REMAT_GROUP = group + _GTP_WEIGHT_REMAT_GLOBAL_RANKS = gtp_ranks + + # Disable Gloo under GTP_remat (out of scope; the GTP_remat optimizer uses DCP). + if gtp_remat_size > 1: + create_gloo_process_groups = False + # Set NCCL_COLLNET_ENABLE to 1 to enable SHARP for the dp group. if sharp_enabled_group == "dp": os.environ["NCCL_COLLNET_ENABLE"] = "1" @@ -841,7 +959,7 @@ def initialize_model_parallel( # is eligible for using the NCCL COLLNET feature. # Therefore, dp-cp group, which potentially requires SHARP-enablement, # need to be created before all the other groups - for ranks_with_cp in decoder_rank_generator.get_ranks('dp-cp'): + for ranks_with_cp in decoder_rank_generator.get_ranks("dp-cp"): group_with_cp = create_group( ranks_with_cp, timeout=timeout, @@ -947,7 +1065,7 @@ def initialize_model_parallel( torch.distributed.barrier(group=group, device_ids=[torch.cuda.current_device()]) torch.cuda.synchronize() - for ranks in decoder_rank_generator.get_ranks('dp'): + for ranks in decoder_rank_generator.get_ranks("dp"): group = create_group( ranks, timeout=timeout, @@ -965,6 +1083,48 @@ def initialize_model_parallel( _DATA_PARALLEL_GROUP_GLOO = group_gloo _DATA_PARALLEL_GLOBAL_RANKS = ranks + # Full data-distribution groups: span gtp_remat explicitly + # ('gtp_remat-dp' / 'gtp_remat-dp-cp'). Used only for batch split, num-microbatches, gradient + # scaling, and reductions covering every distinct-data rank. No Gloo (data distribution uses + # ranks/sizes only). When GTP_remat is inactive they alias the default groups built above. + global _DATA_PARALLEL_GROUP_WITH_GTP_REMAT + global _DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT + global _INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT + if gtp_remat_size > 1: + # Every rank iterates all groups so each create_group collective is entered by all ranks. + for dp_ranks in decoder_rank_generator.get_ranks("gtp_remat-dp"): + group = create_group( + dp_ranks, + timeout=timeout, + pg_options=get_nccl_options("gtp_remat_dp", nccl_comm_cfgs), + group_desc="DATA_PARALLEL_GROUP_WITH_GTP_REMAT", + ) + if rank in dp_ranks: + _DATA_PARALLEL_GROUP_WITH_GTP_REMAT = group + + for dp_cp_ranks in decoder_rank_generator.get_ranks("gtp_remat-dp-cp"): + group = create_group( + dp_cp_ranks, + timeout=timeout, + pg_options=get_nccl_options("gtp_remat_dp_cp", nccl_comm_cfgs), + group_desc="DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT", + ) + if rank in dp_cp_ranks: + _DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT = group + + # GTP_remat requires a single distributed-optimizer instance (asserted above), so the + # per-instance partial full group is just the full group. + _INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT = ( + _DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT + ) + else: + # GTP_remat inactive: the full data-distribution groups coincide with the defaults. + _DATA_PARALLEL_GROUP_WITH_GTP_REMAT = _DATA_PARALLEL_GROUP + _DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT = _DATA_PARALLEL_GROUP_WITH_CP + _INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT = ( + _INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP + ) + # Build the context-parallel groups. global _CONTEXT_PARALLEL_GROUP global _CONTEXT_PARALLEL_GLOBAL_RANKS @@ -994,11 +1154,12 @@ def initialize_model_parallel( if rank in ranks: _HIERARCHICAL_CONTEXT_PARALLEL_GROUPS = hierarchical_groups - # Build the model-parallel groups. + # Model-parallel groups (TP × GTP_remat × PP). gtp_remat is a RankGenerator axis, so the + # 'tp-gtp_remat-pp' token spans it directly; with gtp_remat=1 it reduces to plain tp-pp groups. global _MODEL_PARALLEL_GROUP global _MODEL_PARALLEL_GLOBAL_RANKS assert _MODEL_PARALLEL_GROUP is None, 'model parallel group is already initialized' - for ranks in decoder_rank_generator.get_ranks('tp-pp'): + for ranks in decoder_rank_generator.get_ranks('tp-gtp_remat-pp'): group = create_group( ranks, timeout=timeout, @@ -1148,7 +1309,10 @@ def initialize_model_parallel( assert ( _TENSOR_AND_DATA_PARALLEL_GROUP is None ), 'Tensor + data parallel group is already initialized' - for ranks in decoder_rank_generator.get_ranks('tp-dp-cp'): + # Spans gtp_remat (like dp): gtp_remat peers are distinct-data ranks, so this group serves both + # FP8 amax reduction and the MoE router's expert-bias / load-balancing token reduction. The + # gtp_remat axis is a no-op when its size is 1. + for ranks in decoder_rank_generator.get_ranks('tp-gtp_remat-dp-cp'): group = create_group( ranks, timeout=timeout, @@ -1157,7 +1321,7 @@ def initialize_model_parallel( ) if rank in ranks: _TENSOR_AND_DATA_PARALLEL_GROUP_WITH_CP = group - for ranks in decoder_rank_generator.get_ranks('tp-dp'): + for ranks in decoder_rank_generator.get_ranks('tp-gtp_remat-dp'): group = create_group( ranks, timeout=timeout, @@ -1182,6 +1346,26 @@ def initialize_model_parallel( _TENSOR_AND_CONTEXT_PARALLEL_GROUP = group ### Expert-related parallel groups initialization + # Build the expert generalized tensor parallel group + # Expert GTP_remat overlaps with the expert DP domain (experts don't use CP). + global _EXPERT_GTP_WEIGHT_REMAT_GROUP + global _EXPERT_GTP_WEIGHT_REMAT_GLOBAL_RANKS + assert ( + _EXPERT_GTP_WEIGHT_REMAT_GROUP is None + ), 'Expert generalized tensor parallel group is already initialized' + # EGTP shard groups are get_ranks('gtp_remat') on the expert generator (singletons when + # expert_gtp_remat_size == 1). See RankGenerator.get_gtp_ranks. + for egtp_ranks in expert_decoder_rank_generator.get_gtp_ranks(expert_gtp_remat_size): + group = create_group( + egtp_ranks, + timeout=timeout, + pg_options=get_nccl_options("expt_gtp_remat", nccl_comm_cfgs), + group_desc="EXPERT_GTP_WEIGHT_REMAT_GROUP", + ) + if rank in egtp_ranks: + _EXPERT_GTP_WEIGHT_REMAT_GROUP = group + _EXPERT_GTP_WEIGHT_REMAT_GLOBAL_RANKS = egtp_ranks + # Build the expert model parallel group global _EXPERT_MODEL_PARALLEL_GROUP, _EXPERT_MODEL_PARALLEL_RANKS assert _EXPERT_MODEL_PARALLEL_GROUP is None, 'Expert parallel group is already initialized' @@ -1241,6 +1425,22 @@ def initialize_model_parallel( if rank in ranks: _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP = group + # Expert+tensor+pipeline group merged across EGTP peers — expert analog of the dense + # _MODEL_PARALLEL_GROUP merge (above). The 'tp-ep-gtp_remat-pp' token spans the egtp axis; with + # expert_gtp_remat_size=1 it reduces to the plain tp-ep-pp groups. Merging gives EGTP peers + # distinct ranks; see docs/api-guide/core/generalized_tensor_parallel.md §3.3 + # (Optimizer state) for the DCP-collision rationale. + global _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP_WITH_EGTP + for ranks in expert_decoder_rank_generator.get_ranks('tp-ep-gtp_remat-pp'): + group = create_group( + ranks, + timeout=timeout, + pg_options=get_nccl_options("tp_ep_gtp_remat_pp", nccl_comm_cfgs), + group_desc="EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP_WITH_EGTP", + ) + if rank in ranks: + _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP_WITH_EGTP = group + # Build the expert data parallel group global _EXPERT_DATA_PARALLEL_GROUP assert _EXPERT_DATA_PARALLEL_GROUP is None, "Expert data group is already initialized" @@ -1266,7 +1466,10 @@ def initialize_model_parallel( expert_data_parallel_size // num_distributed_optimizer_instances ) - for ranks in expert_decoder_rank_generator.get_ranks('dp'): + # Gloo only on the non-EGTP path (EGTP + Gloo out of scope; the EGTP optimizer uses DCP). + if expert_gtp_remat_size > 1: + create_gloo_process_groups = False + for ranks in expert_decoder_rank_generator.get_ranks("dp"): group = create_group( ranks, timeout=timeout, @@ -1322,6 +1525,29 @@ def initialize_model_parallel( else: _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP = _EXPERT_DATA_PARALLEL_GROUP _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP_GLOO = _EXPERT_DATA_PARALLEL_GROUP_GLOO + # Full expert data-distribution group: spans gtp_remat explicitly. Used only + # where distinct-data distribution matters; no Gloo. Aliases the default when EGTP is inactive. + global _EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT + global _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT + if expert_gtp_remat_size > 1: + for dp_ranks in expert_decoder_rank_generator.get_ranks("gtp_remat-dp"): + group = create_group( + dp_ranks, + timeout=timeout, + pg_options=get_nccl_options("ep_gtp_remat_dp", nccl_comm_cfgs), + group_desc="EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT", + ) + if rank in dp_ranks: + _EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT = group + _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT = ( + _EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT + ) + else: + _EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT = _EXPERT_DATA_PARALLEL_GROUP + _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT = ( + _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP + ) + ### End of expert related parallel groups initialization # build the intra distributed optimizer instance group @@ -1330,21 +1556,40 @@ def initialize_model_parallel( _INTRA_DISTRIBUTED_OPTIMIZER_INSTANCE_GROUP is None ), "Intra distributed optimizer instance group is already initialized" - model_parallel_group_id = 0 - intra_dist_opt_ranks = [] - for ranks in expert_decoder_rank_generator.get_ranks('tp-ep-pp'): - model_parallel_group_id += 1 - intra_dist_opt_ranks.extend(ranks) - if model_parallel_group_id % intra_partial_expert_data_parallel_size == 0: - intra_dist_opt_instance_group = create_group( - intra_dist_opt_ranks, - timeout=timeout, - pg_options=get_nccl_options("intra_dist_opt_instance", nccl_comm_cfgs), - group_desc="INTRA_DISTRIBUTED_OPTIMIZER_INSTANCE_GROUP", - ) - if rank in intra_dist_opt_ranks: - _INTRA_DISTRIBUTED_OPTIMIZER_INSTANCE_GROUP = intra_dist_opt_instance_group - intra_dist_opt_ranks = [] + if gtp_remat_size > 1 or expert_gtp_remat_size > 1: + # GTP_remat requires num_distributed_optimizer_instances == 1 (asserted above); dist-opt + # grad-stats group (used only for grad-norm + num_zeros reductions) must span the ENTIRE + # world. The per-instance accumulation below would NOT: gtp/egtp are factored out of + # expert_data_parallel_size (via expert_gtp_remat_size), so expert-generator groups omit + # gtp/egtp axes — under-counting the grad-norm for gtp/egtp-sharded params. Build one + # full-world group from all tp-ep-pp groups instead (get_ranks already applies rank_offset). + all_ranks = sorted( + r for ranks in expert_decoder_rank_generator.get_ranks('tp-ep-pp') for r in ranks + ) + intra_dist_opt_instance_group = create_group( + all_ranks, + timeout=timeout, + pg_options=get_nccl_options("intra_dist_opt_instance", nccl_comm_cfgs), + group_desc="INTRA_DISTRIBUTED_OPTIMIZER_INSTANCE_GROUP", + ) + if rank in all_ranks: + _INTRA_DISTRIBUTED_OPTIMIZER_INSTANCE_GROUP = intra_dist_opt_instance_group + else: + model_parallel_group_id = 0 + intra_dist_opt_ranks = [] + for ranks in expert_decoder_rank_generator.get_ranks('tp-ep-pp'): + model_parallel_group_id += 1 + intra_dist_opt_ranks.extend(ranks) + if model_parallel_group_id % intra_partial_expert_data_parallel_size == 0: + intra_dist_opt_instance_group = create_group( + intra_dist_opt_ranks, + timeout=timeout, + pg_options=get_nccl_options("intra_dist_opt_instance", nccl_comm_cfgs), + group_desc="INTRA_DISTRIBUTED_OPTIMIZER_INSTANCE_GROUP", + ) + if rank in intra_dist_opt_ranks: + _INTRA_DISTRIBUTED_OPTIMIZER_INSTANCE_GROUP = intra_dist_opt_instance_group + intra_dist_opt_ranks = [] # Initialize global memory buffer # This isn't really "parallel state" but there isn't another good place to @@ -1392,11 +1637,19 @@ def create_all_gather_groups(for_expert_parallelism=False, timeout=None, nccl_co tp_size = get_tensor_model_parallel_world_size() ep_size = get_expert_model_parallel_world_size() dp_size = get_data_parallel_world_size() + gtp_remat_size = get_gtp_weight_remat_world_size() or 1 # Create regular DP all-gather group dp_cp_ag_group = None decoder_rank_gen = RankGenerator( - tp=tp_size, ep=1, dp=dp_size, pp=pp_size, cp=cp_size, order='tp-cp-ep-dp-pp', rank_offset=0 + tp=tp_size, + ep=1, + dp=dp_size, + pp=pp_size, + cp=cp_size, + gtp_remat=gtp_remat_size, + order=_inject_gtp_remat_axis('tp-cp-ep-dp-pp', after='tp'), + rank_offset=0, ) for ranks_with_cp in decoder_rank_gen.get_ranks('dp-cp'): @@ -1414,6 +1667,7 @@ def create_all_gather_groups(for_expert_parallelism=False, timeout=None, nccl_co if for_expert_parallelism and ep_size > 1: expert_tp_size = get_expert_tensor_parallel_world_size() expert_dp_size = get_expert_data_parallel_world_size() + egtp_remat_size = get_expert_gtp_weight_remat_world_size() or 1 expert_rank_gen = RankGenerator( tp=expert_tp_size, @@ -1421,7 +1675,8 @@ def create_all_gather_groups(for_expert_parallelism=False, timeout=None, nccl_co dp=expert_dp_size, pp=pp_size, cp=1, - order='tp-cp-ep-dp-pp', + gtp_remat=egtp_remat_size, + order=_inject_gtp_remat_axis('tp-cp-ep-dp-pp', after='ep'), rank_offset=0, ) @@ -1470,6 +1725,42 @@ def get_tensor_model_parallel_group(check_initialized=True): return _TENSOR_MODEL_PARALLEL_GROUP +def get_gtp_weight_remat_group(check_initialized=True): + """Get the parameter-sharding group the caller rank belongs to.""" + if check_initialized: + assert ( + _GTP_WEIGHT_REMAT_GROUP is not None + ), "generalized tensor parallel group is not initialized" + return _GTP_WEIGHT_REMAT_GROUP + + +def get_gtp_weight_remat_world_size(): + """Return world size for the parameter-sharding group.""" + if torch.distributed.is_available() and torch.distributed.is_initialized(): + group = get_gtp_weight_remat_group(check_initialized=False) + return group.size() if group is not None else 0 + else: + return 0 + + +def get_gtp_weight_remat_rank(): + """Return caller's rank in the parameter-sharding group.""" + if torch.distributed.is_available() and torch.distributed.is_initialized(): + group = get_gtp_weight_remat_group(check_initialized=False) + return group.rank() if group is not None else 0 + else: + return 0 + + +def get_gtp_weight_remat_global_ranks(check_initialized=True): + """Get all global ranks of the parameter-sharding group that the caller rank belongs to.""" + if check_initialized: + assert ( + _GTP_WEIGHT_REMAT_GLOBAL_RANKS is not None + ), "generalized tensor parallel group is not initialized" + return _GTP_WEIGHT_REMAT_GLOBAL_RANKS + + def get_pipeline_model_parallel_group(check_initialized=True): """Get the pipeline-model-parallel group the caller rank belongs to.""" if check_initialized: @@ -1479,22 +1770,56 @@ def get_pipeline_model_parallel_group(check_initialized=True): return _PIPELINE_MODEL_PARALLEL_GROUP -def get_data_parallel_group(with_context_parallel=False, partial_data_parallel=False): - """Get the data-parallel group the caller rank belongs to.""" - if with_context_parallel: - if partial_data_parallel: - assert ( - _INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP is not None - ), "Intra partial data parallel group is not initialized" - return _INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP - assert ( - _DATA_PARALLEL_GROUP_WITH_CP is not None - ), "data parallel group with context parallel combined is not initialized" - return _DATA_PARALLEL_GROUP_WITH_CP +def get_data_parallel_group( + with_context_parallel=False, with_gtp_remat=True, partial_data_parallel=False +): + """Get the data-parallel group the caller rank belongs to. + + GTP_remat is an independent axis layered on DP. + DEFAULT (``with_gtp_remat=True``): full data-distribution group (replicate_DP x + gtp_remat) — gtp_remat peers hold distinct micro-batches, so use it for batch split, + num-microbatches, grad scaling, and reductions over all distinct-data ranks. + ``with_gtp_remat=False``: replicate group — grad all-reduce, optimizer-state + sharding, checkpoint replicas. + + Args: + with_context_parallel: If True, include context-parallel ranks. + with_gtp_remat: True (default) = full data-distribution group; False = replicate. + partial_data_parallel: If True, return partial DP group (requires with_context_parallel). + """ + assert ( + with_context_parallel or not partial_data_parallel + ), "Partial DP for Optimizer needs to include CP" + # (with_cp, partial_data_parallel) -> (group, description). Globals are read at call time + # (assigned during initialize_model_parallel). partial requires CP, so the (False, True) row + # is unreachable and omitted. + if with_gtp_remat: + group_table = { + (False, False): ( + _DATA_PARALLEL_GROUP_WITH_GTP_REMAT, + "data parallel group (with GTP_remat)", + ), + (True, False): ( + _DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT, + "data parallel group with CP (with GTP_remat)", + ), + (True, True): ( + _INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT, + "intra partial data parallel group with CP (with GTP_remat)", + ), + } else: - assert _DATA_PARALLEL_GROUP is not None, "data parallel group is not initialized" - assert partial_data_parallel == False, "Partial DP for Optimizer needs to include CP" - return _DATA_PARALLEL_GROUP + group_table = { + (False, False): (_DATA_PARALLEL_GROUP, "data parallel group"), + (True, False): (_DATA_PARALLEL_GROUP_WITH_CP, "data parallel group with CP"), + (True, True): ( + _INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP, + "intra partial data parallel group with CP", + ), + } + group, description = group_table[(with_context_parallel, partial_data_parallel)] + assert group is not None, f"{description} is not initialized" + return group def get_data_parallel_group_gloo(with_context_parallel=False, partial_data_parallel=False): @@ -1590,7 +1915,11 @@ def get_amax_reduction_group(with_context_parallel=False, tp_only_amax_red=False def get_tensor_and_data_parallel_group(check_initialized=True, with_context_parallel=False): - """Get the tensor- and data-parallel group the caller rank belongs to.""" + """Get the tensor- and data-parallel group the caller rank belongs to. + + The group spans gtp_remat (like dp), so it serves both FP8 amax reduction and the MoE router's + expert-bias / load-balancing token reduction across every distinct-data rank. + """ if with_context_parallel: if check_initialized: assert ( @@ -1802,14 +2131,22 @@ def get_pipeline_model_parallel_prev_rank(): return _PIPELINE_GLOBAL_RANKS[(rank_in_pipeline - 1) % world_size] -def get_data_parallel_world_size(with_context_parallel=False, partial_data_parallel=False): - """Return world size for the data parallel group.""" +def get_data_parallel_world_size( + with_context_parallel=False, with_gtp_remat=True, partial_data_parallel=False +): + """Return the data-parallel world size. + + DEFAULT (with_gtp_remat=True): full degree (replicate_DP x gtp_remat). + with_gtp_remat=False: replicate degree. + """ global _MPU_DATA_PARALLEL_WORLD_SIZE if _MPU_DATA_PARALLEL_WORLD_SIZE is not None: return _MPU_DATA_PARALLEL_WORLD_SIZE if torch.distributed.is_available() and torch.distributed.is_initialized(): return get_data_parallel_group( - with_context_parallel=with_context_parallel, partial_data_parallel=partial_data_parallel + with_context_parallel=with_context_parallel, + with_gtp_remat=with_gtp_remat, + partial_data_parallel=partial_data_parallel, ).size() else: return 0 @@ -1821,14 +2158,22 @@ def set_data_parallel_rank(rank): _MPU_DATA_PARALLEL_RANK = rank -def get_data_parallel_rank(with_context_parallel=False, partial_data_parallel=False): - """Return caller's rank in the data-parallel group.""" +def get_data_parallel_rank( + with_context_parallel=False, with_gtp_remat=True, partial_data_parallel=False +): + """Return the caller's data-parallel rank. + + DEFAULT (with_gtp_remat=True): rank in the full group (replicate_DP x gtp_remat). + with_gtp_remat=False: rank in the replicate group. + """ global _MPU_DATA_PARALLEL_RANK if _MPU_DATA_PARALLEL_RANK is not None: return _MPU_DATA_PARALLEL_RANK if torch.distributed.is_available() and torch.distributed.is_initialized(): return get_data_parallel_group( - with_context_parallel=with_context_parallel, partial_data_parallel=partial_data_parallel + with_context_parallel=with_context_parallel, + with_gtp_remat=with_gtp_remat, + partial_data_parallel=partial_data_parallel, ).rank() else: return 0 @@ -1867,6 +2212,42 @@ def get_tensor_and_context_parallel_rank(): ### Expert-related parallel states functions +def get_expert_gtp_weight_remat_group(check_initialized=True): + """Get the expert-parameter-sharding group the caller rank belongs to.""" + if check_initialized: + assert ( + _EXPERT_GTP_WEIGHT_REMAT_GROUP is not None + ), "expert generalized tensor parallel group is not initialized" + return _EXPERT_GTP_WEIGHT_REMAT_GROUP + + +def get_expert_gtp_weight_remat_world_size(): + """Return world size for the expert-parameter-sharding group.""" + if torch.distributed.is_available() and torch.distributed.is_initialized(): + group = get_expert_gtp_weight_remat_group(check_initialized=False) + return group.size() if group is not None else 0 + else: + return 0 + + +def get_expert_gtp_weight_remat_rank(): + """Return caller's rank in the expert-parameter-sharding group.""" + if torch.distributed.is_available() and torch.distributed.is_initialized(): + group = get_expert_gtp_weight_remat_group(check_initialized=False) + return group.rank() if group is not None else 0 + else: + return 0 + + +def get_expert_gtp_weight_remat_global_ranks(check_initialized=True): + """Get all global ranks of the expert-parameter-sharding group that the caller rank belongs to.""" + if check_initialized: + assert ( + _EXPERT_GTP_WEIGHT_REMAT_GLOBAL_RANKS is not None + ), "expert generalized tensor parallel group is not initialized" + return _EXPERT_GTP_WEIGHT_REMAT_GLOBAL_RANKS + + def get_expert_model_parallel_group(check_initialized=True): """Get the expert-model-parallel group the caller rank belongs to.""" if check_initialized: @@ -1988,8 +2369,23 @@ def get_expert_tensor_and_model_parallel_rank(): return 0 -def get_expert_tensor_model_pipeline_parallel_group(check_initialized=True): - """Get expert tensor-model-pipeline parallel group.""" +def get_expert_tensor_model_pipeline_parallel_group(check_initialized=True, with_egtp_remat=False): + """Get expert tensor-model-pipeline parallel group. + + Args: + check_initialized: If True (default), asserts the group has been created. + with_egtp_remat: If True, return the EGTP-merged variant — the analog of dense + ``get_model_parallel_group()`` (which merges across GTP peers). Use this when you + need a group whose rank uniquely identifies each (ETP, EP, PP, EGTP) position; + e.g. for the MoE distributed optimizer's ``data_parallel_group_idx``. Identical + to the vanilla group when EGTP_remat_size=1. + """ + if with_egtp_remat: + if check_initialized: + assert ( + _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP_WITH_EGTP is not None + ), "Expert tensor-model-pipeline parallel group with EGTP is not initialized" + return _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP_WITH_EGTP if check_initialized: assert ( _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP is not None @@ -1997,20 +2393,36 @@ def get_expert_tensor_model_pipeline_parallel_group(check_initialized=True): return _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP -def get_expert_data_parallel_group(check_initialized=True, partial_expert_data_parallel=False): - """Get expert data parallel group.""" - if partial_expert_data_parallel: - if check_initialized: - assert ( - _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP is not None - ), "Intra partial expert data parallel group is not initialized" - return _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP - else: - if check_initialized: - assert ( - _EXPERT_DATA_PARALLEL_GROUP is not None - ), "Expert data parallel group is not initialized" - return _EXPERT_DATA_PARALLEL_GROUP +def get_expert_data_parallel_group( + check_initialized=True, with_gtp_remat=True, partial_expert_data_parallel=False +): + """Get the expert data parallel group. + + DEFAULT (with_gtp_remat=True): full group for data distribution (EGTP_remat peers + hold distinct micro-batches). + with_gtp_remat=False: replicate group — expert grad all-reduce, optimizer state, + checkpoint replicas. + """ + # (with_gtp_remat, partial_expert_data_parallel) -> (group, description). Read at call time. + group_table = { + (False, False): (_EXPERT_DATA_PARALLEL_GROUP, "Expert data parallel group"), + (False, True): ( + _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP, + "Intra partial expert data parallel group", + ), + (True, False): ( + _EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT, + "Expert data parallel group (with GTP_remat)", + ), + (True, True): ( + _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT, + "Intra partial expert data parallel group (with GTP_remat)", + ), + } + group, description = group_table[(with_gtp_remat, partial_expert_data_parallel)] + if check_initialized: + assert group is not None, f"{description} is not initialized" + return group def get_expert_data_parallel_group_gloo(partial_expert_data_parallel=False): @@ -2027,21 +2439,21 @@ def get_expert_data_parallel_group_gloo(partial_expert_data_parallel=False): return _EXPERT_DATA_PARALLEL_GROUP_GLOO -def get_expert_data_parallel_rank(partial_expert_data_parallel=False): - """Return caller's rank in the expert data parallel group.""" +def get_expert_data_parallel_rank(with_gtp_remat=True, partial_expert_data_parallel=False): + """Return the caller's expert-data-parallel rank (default: EGTP_remat-inclusive).""" if torch.distributed.is_available() and torch.distributed.is_initialized(): return get_expert_data_parallel_group( - partial_expert_data_parallel=partial_expert_data_parallel + with_gtp_remat=with_gtp_remat, partial_expert_data_parallel=partial_expert_data_parallel ).rank() else: return 0 -def get_expert_data_parallel_world_size(partial_expert_data_parallel=False): - """Return world size for the expert data parallel group.""" +def get_expert_data_parallel_world_size(with_gtp_remat=True, partial_expert_data_parallel=False): + """Return the expert-data-parallel world size (default: EGTP_remat-inclusive).""" if torch.distributed.is_available() and torch.distributed.is_initialized(): return get_expert_data_parallel_group( - partial_expert_data_parallel=partial_expert_data_parallel + with_gtp_remat=with_gtp_remat, partial_expert_data_parallel=partial_expert_data_parallel ).size() else: return 0 @@ -2096,6 +2508,7 @@ def get_all_ranks(): pipeline-model-parallel and expert-model-parallel groups.""" ranks = [ get_tensor_model_parallel_rank(), + get_gtp_weight_remat_rank(), get_data_parallel_rank(), get_context_parallel_rank(), get_pipeline_model_parallel_rank(), @@ -2123,15 +2536,30 @@ def destroy_model_parallel(): global _TENSOR_MODEL_PARALLEL_GROUP _TENSOR_MODEL_PARALLEL_GROUP = None + global _GTP_WEIGHT_REMAT_GROUP + _GTP_WEIGHT_REMAT_GROUP = None + + global _GTP_WEIGHT_REMAT_GLOBAL_RANKS + _GTP_WEIGHT_REMAT_GLOBAL_RANKS = None + global _PIPELINE_MODEL_PARALLEL_GROUP _PIPELINE_MODEL_PARALLEL_GROUP = None global _DATA_PARALLEL_GROUP _DATA_PARALLEL_GROUP = None + global _DATA_PARALLEL_GROUP_WITH_GTP_REMAT + _DATA_PARALLEL_GROUP_WITH_GTP_REMAT = None + global _DATA_PARALLEL_GROUP_WITH_CP _DATA_PARALLEL_GROUP_WITH_CP = None + global _DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT + _DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT = None + + global _INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT + _INTRA_PARTIAL_DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT = None + global _CONTEXT_PARALLEL_GROUP _CONTEXT_PARALLEL_GROUP = None @@ -2201,6 +2629,12 @@ def destroy_model_parallel(): _DATA_PARALLEL_GROUP_WITH_CP_GLOO = None # Destroy parallel state related to expert parallelism. + global _EXPERT_GTP_WEIGHT_REMAT_GROUP + _EXPERT_GTP_WEIGHT_REMAT_GROUP = None + + global _EXPERT_GTP_WEIGHT_REMAT_GLOBAL_RANKS + _EXPERT_GTP_WEIGHT_REMAT_GLOBAL_RANKS = None + global _EXPERT_MODEL_PARALLEL_GROUP _EXPERT_MODEL_PARALLEL_GROUP = None @@ -2225,9 +2659,18 @@ def destroy_model_parallel(): global _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP = None + global _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP_WITH_EGTP + _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP_WITH_EGTP = None + global _EXPERT_DATA_PARALLEL_GROUP _EXPERT_DATA_PARALLEL_GROUP = None + global _EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT + _EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT = None + + global _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT + _INTRA_PARTIAL_EXPERT_DATA_PARALLEL_GROUP_WITH_GTP_REMAT = None + global _EXPERT_DATA_PARALLEL_GROUP_GLOO if ( _EXPERT_DATA_PARALLEL_GROUP_GLOO is not None diff --git a/megatron/core/pipeline_parallel/combined_1f1b.py b/megatron/core/pipeline_parallel/combined_1f1b.py index 81524363993..206a50d76ad 100644 --- a/megatron/core/pipeline_parallel/combined_1f1b.py +++ b/megatron/core/pipeline_parallel/combined_1f1b.py @@ -381,12 +381,6 @@ def forward_backward_step(): unwrapped_model = get_attr_wrapped_model( f_model, "build_schedule_plan", return_model_obj=True ) - from megatron.core.models.gpt.gpt_model import GPTModel - - assert isinstance(unwrapped_model, GPTModel), ( - "The final unwrapped model must be a GPTModel instance " - "since only GPTModel is supported for EP A2A overlapping." - ) f_schedule_plan, loss_func = forward_step_func( data_iterator, unwrapped_model, return_schedule_plan=True ) diff --git a/megatron/core/pipeline_parallel/schedules.py b/megatron/core/pipeline_parallel/schedules.py index 897d3d2516b..4eeb898a552 100644 --- a/megatron/core/pipeline_parallel/schedules.py +++ b/megatron/core/pipeline_parallel/schedules.py @@ -671,6 +671,30 @@ def check_first_val_step(first_val_step, forward_only, cond): return cond +def _build_default_pg_collection() -> ProcessGroupCollection: + """Build a ``ProcessGroupCollection`` from the global ``parallel_state`` defaults. + + Used by the schedule entry points as the fallback when the caller does not + supply a ``pg_collection`` explicitly. + """ + pg_collection = ProcessGroupCollection() + pg_collection.tp = parallel_state.get_tensor_model_parallel_group() + pg_collection.cp = parallel_state.get_context_parallel_group() + pg_collection.embd = parallel_state.get_embedding_group(check_initialized=False) + pg_collection.pos_embd = parallel_state.get_position_embedding_group(check_initialized=False) + pg_collection.pp = parallel_state.get_pipeline_model_parallel_group() + pg_collection.dp_cp = parallel_state.get_data_parallel_group( + with_context_parallel=True, partial_data_parallel=False + ) + pg_collection.tp_dp_cp = parallel_state.get_tensor_and_data_parallel_group( + with_context_parallel=True + ) + pg_collection.dp = parallel_state.get_data_parallel_group( + with_context_parallel=False, partial_data_parallel=False + ) + return pg_collection + + def forward_backward_no_pipelining( *, forward_step_func, @@ -691,23 +715,7 @@ def forward_backward_no_pipelining( """Run forward and backward passes with no pipeline parallelism""" if pg_collection is None: - tp_group = parallel_state.get_tensor_model_parallel_group() - cp_group = parallel_state.get_context_parallel_group() - embd_group = parallel_state.get_embedding_group(check_initialized=False) - pp_group = parallel_state.get_pipeline_model_parallel_group() - pos_emb_group = parallel_state.get_position_embedding_group(check_initialized=False) - pg_collection = ProcessGroupCollection() - pg_collection.tp = tp_group - pg_collection.cp = cp_group - pg_collection.embd = embd_group - pg_collection.pos_embd = pos_emb_group - pg_collection.pp = pp_group - pg_collection.dp_cp = parallel_state.get_data_parallel_group( - with_context_parallel=True, partial_data_parallel=False - ) - pg_collection.tp_dp_cp = parallel_state.get_tensor_and_data_parallel_group( - with_context_parallel=True - ) + pg_collection = _build_default_pg_collection() elif pg_collection is not None: assert hasattr(pg_collection, 'tp'), "pg_collection must have tp" @@ -1131,25 +1139,10 @@ def forward_backward_pipelining_with_interleaving( p2p_communicator = P2PCommunicator( pp_group=parallel_state.get_pipeline_model_parallel_group(), config=config ) - tp_group = parallel_state.get_tensor_model_parallel_group() - cp_group = parallel_state.get_context_parallel_group() + pg_collection = _build_default_pg_collection() + tp_group = pg_collection.tp + cp_group = pg_collection.cp cp_size = cp_group.size() - embd_group = parallel_state.get_embedding_group(check_initialized=False) - pp_group = parallel_state.get_pipeline_model_parallel_group() - pos_emb_group = parallel_state.get_position_embedding_group(check_initialized=False) - - pg_collection = ProcessGroupCollection() - pg_collection.tp = tp_group - pg_collection.cp = cp_group - pg_collection.embd = embd_group - pg_collection.pos_embd = pos_emb_group - pg_collection.pp = pp_group - pg_collection.dp_cp = parallel_state.get_data_parallel_group( - with_context_parallel=True, partial_data_parallel=False - ) - pg_collection.tp_dp_cp = parallel_state.get_tensor_and_data_parallel_group( - with_context_parallel=True - ) elif p2p_communicator is not None and pg_collection is not None: model_type = get_model_type(model[0]) @@ -2328,25 +2321,10 @@ def forward_backward_pipelining_without_interleaving( p2p_communicator = P2PCommunicator( pp_group=parallel_state.get_pipeline_model_parallel_group(), config=config ) - tp_group = parallel_state.get_tensor_model_parallel_group() - cp_group = parallel_state.get_context_parallel_group() + pg_collection = _build_default_pg_collection() + tp_group = pg_collection.tp + cp_group = pg_collection.cp cp_size = cp_group.size() - embd_group = parallel_state.get_embedding_group(check_initialized=False) - pos_emb_group = parallel_state.get_position_embedding_group(check_initialized=False) - pp_group = parallel_state.get_pipeline_model_parallel_group() - - pg_collection = ProcessGroupCollection() - pg_collection.tp = tp_group - pg_collection.pp = pp_group - pg_collection.embd = embd_group - pg_collection.pos_embd = pos_emb_group - pg_collection.cp = cp_group - pg_collection.dp_cp = parallel_state.get_data_parallel_group( - with_context_parallel=True, partial_data_parallel=False - ) - pg_collection.tp_dp_cp = parallel_state.get_tensor_and_data_parallel_group( - with_context_parallel=True - ) elif p2p_communicator is not None and pg_collection is not None: assert hasattr(p2p_communicator, 'config'), "p2p_communicator must have a config" diff --git a/megatron/core/pipeline_parallel/utils.py b/megatron/core/pipeline_parallel/utils.py index e5afd5ac388..8ede746bf9b 100644 --- a/megatron/core/pipeline_parallel/utils.py +++ b/megatron/core/pipeline_parallel/utils.py @@ -17,9 +17,46 @@ nvtx_range_push, ) +try: + from transformer_engine.pytorch.ep import is_symm_backed +except ImportError: + is_symm_backed = None + logger = logging.getLogger(__name__) +class StageDispatchBwdGrad(torch.autograd.Function): + """1F1B + NCCL-EP zero-copy only: redirect the dispatch-backward grad into the persistent + symm buffer so the one-sided ``dispatch_bwd`` can consume it. + + Under the 1F1B overlap schedule the dispatch output is consumed by the next node, which + detaches it into a leaf; autograd therefore hands ``dispatch_bwd`` a non-symm + ``AccumulateGrad`` clone. Applying this identity node to the dispatch output — while it is + still inside the dispatch node's own graph segment — makes it the sole consumer, moving that + accumulation to *our* output; the backward then does a single plain->symm copy into the + dispatcher's ``_zc_bwd_token_buf``. That buffer is free to stage into precisely because + ``get_expert_zero_copy_buffers`` withholds it from the op-fuser under overlap. + Forward is identity (no numeric effect). + """ + + @staticmethod + def forward(ctx, dispatched_tokens, token_dispatcher): # type: ignore[override] + """Identity forward; stashes the dispatcher so backward can reach its symm buffer.""" + ctx.token_dispatcher = token_dispatcher + return dispatched_tokens + + @staticmethod + def backward(ctx, grad): # type: ignore[override] + """Stage the incoming gradient into the symm dispatch-backward buffer.""" + buf = ctx.token_dispatcher._comm_manager._zc_bwd_token_buf + assert buf is not None, "zero-copy staging buffer not allocated before dispatch-backward" + assert ( + buf.shape == grad.shape + ), f"dispatch-bwd grad {tuple(grad.shape)} != staging buffer {tuple(buf.shape)}" + buf.copy_(grad) + return buf, None + + def is_pp_first_stage(pp_group: torch.distributed.ProcessGroup): """Return True if in the first pipeline model-parallel stage, False otherwise.""" return get_pg_rank(pp_group) == 0 @@ -158,6 +195,7 @@ def __init__( name: str = "schedule_node", forward_nvtx_name: Optional[str] = None, backward_nvtx_name: Optional[str] = None, + ncclep_zero_copy: bool = False, ): """Initialize a schedule node. @@ -186,6 +224,7 @@ def __init__( self.stream = stream self.event = event self.free_input = free_input + self.ncclep_zero_copy = ncclep_zero_copy self.inputs = None self.outputs = None @@ -236,7 +275,13 @@ def _forward(self, *inputs): for input in inputs: if input is not None: input.record_stream(self.stream) - input.untyped_storage().resize_(0) + # Skip symmetric-memory (zero-copy EP) buffers + if not ( + self.ncclep_zero_copy + and is_symm_backed is not None + and is_symm_backed(input) + ): + input.untyped_storage().resize_(0) return self.output diff --git a/megatron/core/post_training/modelopt/layers.py b/megatron/core/post_training/modelopt/layers.py index 04e03a36458..5f1a746e95b 100644 --- a/megatron/core/post_training/modelopt/layers.py +++ b/megatron/core/post_training/modelopt/layers.py @@ -159,7 +159,12 @@ def __init__( for param in self.parameters(): if is_expert: # Reduce the gradient on the expert_data_parallel group for expert linear layers - setattr(param, "allreduce", self.config.expert_model_parallel_size == 1) + use_expert_groups = ( + self.config.expert_model_parallel_size > 1 + or self.config.expert_tensor_parallel_size + != self.config.tensor_model_parallel_size + ) + setattr(param, "allreduce", not use_expert_groups) else: # Reduce the gradient on DP group setattr(param, "allreduce", True) diff --git a/megatron/core/process_groups_config.py b/megatron/core/process_groups_config.py index 6c1e3651387..ccb6dce0eb8 100644 --- a/megatron/core/process_groups_config.py +++ b/megatron/core/process_groups_config.py @@ -44,11 +44,17 @@ class ProcessGroupCollection: expt_tp: Expert tensor parallel group tp_ep: Tensor and expert parallel group tp_ep_pp: Tensor, expert, and pipeline parallel group + tp_ep_pp_with_egtp_remat: tp_ep_pp merged across EGTP peers (dense ``mp`` analog); + identical to ``tp_ep_pp`` when EGTP_remat_size=1 # Data Parallelism Groups dp: Data parallel process group dp_cp: Data and context parallel group + dp_cp_gtp_remat: Full data-distribution group, dp_cp x gtp_remat; + identical to dp_cp when GTP_remat_size=1 expt_dp: Expert data parallel group + expt_dp_gtp_remat: Full expert data-distribution group, expt_dp x egtp_remat; + identical to expt_dp when EGTP_remat_size=1 intra_dp_cp: Intra partial data parallel group intra_expt_dp: Intra partial expert data parallel group inter_dist_opt: Inter distributed optimizer instance group @@ -104,7 +110,12 @@ class ProcessGroupCollection: # _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP tp_ep_pp: torch.distributed.ProcessGroup = field(init=False) - # _TENSOR_AND_DATA_PARALLEL_GROUP_WITH_CP + # _EXPERT_TENSOR_MODEL_PIPELINE_PARALLEL_GROUP_WITH_EGTP — expert "model parallel" group + # merged across EGTP peers (analog of dense ``mp`` under GTP). Identical to ``tp_ep_pp`` + # when EGTP_remat_size=1. + tp_ep_pp_with_egtp_remat: torch.distributed.ProcessGroup = field(init=False) + + # _TENSOR_AND_DATA_PARALLEL_GROUP_WITH_CP (spans gtp_remat, like dp) tp_dp_cp: torch.distributed.ProcessGroup = field(init=False) # Data Parallelism Process Groups @@ -114,9 +125,21 @@ class ProcessGroupCollection: # _DATA_PARALLEL_GROUP_WITH_CP dp_cp: torch.distributed.ProcessGroup = field(init=False) + # _DATA_PARALLEL_GROUP_WITH_CP_WITH_GTP_REMAT — the full data-distribution group, DP x CP x + # gtp_remat. Used where data is split/aggregated across every rank that holds a distinct + # micro-batch (batch split, num_microbatches, loss/metric reductions, non-gtp param broadcast). + # Identical to ``dp_cp`` when gtp_remat_size=1. + dp_cp_gtp_remat: torch.distributed.ProcessGroup = field(init=False) + # Separate dp_cp communicator for param all-gather (AG/RS overlap) dp_cp_ag: torch.distributed.ProcessGroup = field(init=False) + # _GTP_WEIGHT_REMAT_GROUP + gtp_remat: torch.distributed.ProcessGroup = field(init=False) + + # _EXPERT_GTP_WEIGHT_REMAT_GROUP + expt_gtp_remat: torch.distributed.ProcessGroup = field(init=False) + # MoE layers need expt_dp group for sharded state dict # we need this workaround until distributed checkpoint is refactored # to have sharded_state_dict can take the PG and pass it down @@ -124,6 +147,11 @@ class ProcessGroupCollection: # _EXPERT_DATA_PARALLEL_GROUP expt_dp: torch.distributed.ProcessGroup = field(init=False) + # _EXPERT_DATA_PARALLEL_GROUP_WITH_EGTP — full expert data-distribution group, expt_dp x + # egtp_remat. Used for non-egtp expert param broadcast / sharded state dict over the full + # axis. Identical to ``expt_dp`` when egtp_remat_size=1. + expt_dp_gtp_remat: torch.distributed.ProcessGroup = field(init=False) + # _EXPERT_DATA_PARALLEL_GROUP_AG expt_dp_ag: torch.distributed.ProcessGroup = field(init=False) @@ -146,19 +174,27 @@ def __init__(self, **kwargs): else: raise ValueError(f"Unknown attribute: {key}") + def __getattr__(self, name: str): + # Return None for any declared field that was not set during partial construction + # (e.g. when use_mpu_process_groups is called with a subset of required_pgs). + if name in {f.name for f in fields(self.__class__)}: + return None + raise AttributeError(f"'ProcessGroupCollection' object has no attribute '{name}'") + def __repr__(self): """Return a concise representation showing which process groups exist and their sizes.""" active_pgs = [] for field_info in fields(self): - if hasattr(self, field_info.name): - pg = getattr(self, field_info.name) - if pg is None: - active_pgs.append(f"{field_info.name}(None)") - elif isinstance(pg, list): - sizes = [g.size() for g in pg] - active_pgs.append(f"{field_info.name}({sizes})") - else: - active_pgs.append(f"{field_info.name}({pg.size()})") + if field_info.name not in vars(self): + continue + pg = getattr(self, field_info.name) + if pg is None: + continue + elif isinstance(pg, list): + sizes = [g.size() for g in pg] + active_pgs.append(f"{field_info.name}({sizes})") + else: + active_pgs.append(f"{field_info.name}({pg.size()})") return ( f"ProcessGroupCollection({', '.join(active_pgs)})" if active_pgs @@ -212,21 +248,35 @@ def use_mpu_process_groups(cls, required_pgs: Optional[List[str]] = None): parallel_state.get_expert_tensor_model_pipeline_parallel_group, check_initialized=False, ), + 'tp_ep_pp_with_egtp_remat': partial( + parallel_state.get_expert_tensor_model_pipeline_parallel_group, + check_initialized=False, + with_egtp_remat=True, + ), 'embd': partial(parallel_state.get_embedding_group, check_initialized=False), 'pos_embd': partial( parallel_state.get_position_embedding_group, check_initialized=False ), - 'dp': parallel_state.get_data_parallel_group, - 'dp_cp': partial(parallel_state.get_data_parallel_group, with_context_parallel=True), + 'dp': partial(parallel_state.get_data_parallel_group, with_gtp_remat=False), + 'dp_cp': partial( + parallel_state.get_data_parallel_group, + with_context_parallel=True, + with_gtp_remat=False, + ), + 'dp_cp_gtp_remat': partial( + parallel_state.get_data_parallel_group, with_context_parallel=True + ), 'dp_cp_ag': lambda: None, 'intra_dp_cp': partial( parallel_state.get_data_parallel_group, with_context_parallel=True, + with_gtp_remat=False, partial_data_parallel=True, ), 'intra_expt_dp': partial( parallel_state.get_expert_data_parallel_group, check_initialized=False, + with_gtp_remat=False, partial_expert_data_parallel=True, ), 'inter_dist_opt': partial( @@ -239,6 +289,11 @@ def use_mpu_process_groups(cls, required_pgs: Optional[List[str]] = None): ), # TODO (Hepteract): remove this once distributed checkpoint is refactored 'expt_dp': partial( + parallel_state.get_expert_data_parallel_group, + check_initialized=False, + with_gtp_remat=False, + ), + 'expt_dp_gtp_remat': partial( parallel_state.get_expert_data_parallel_group, check_initialized=False ), 'expt_dp_ag': lambda: None, @@ -247,6 +302,12 @@ def use_mpu_process_groups(cls, required_pgs: Optional[List[str]] = None): check_initialized=False, with_context_parallel=True, ), + 'gtp_remat': partial( + parallel_state.get_gtp_weight_remat_group, check_initialized=False + ), + 'expt_gtp_remat': partial( + parallel_state.get_expert_gtp_weight_remat_group, check_initialized=False + ), } assert all( @@ -259,6 +320,19 @@ def use_mpu_process_groups(cls, required_pgs: Optional[List[str]] = None): return cls(**init_dict) + @staticmethod + def is_gtp_remat_active(process_group_dict: Dict) -> bool: + """True iff GTP_remat or EGTP_remat is active (a weight-shard group spans >1 rank). + + Reads 'gtp_remat_group'/'expt_gtp_remat_group' from setup_process_groups_for_* + builders; a None group means that axis is unused. + """ + gtp_remat = process_group_dict.get('gtp_remat_group') + expt_gtp_remat = process_group_dict.get('expt_gtp_remat_group') + return (gtp_remat is not None and gtp_remat.size() > 1) or ( + expt_gtp_remat is not None and expt_gtp_remat.size() > 1 + ) + @staticmethod def setup_process_groups_for_optimizer( pg_collection: Optional['ProcessGroupCollection'], @@ -294,22 +368,30 @@ def setup_process_groups_for_optimizer( if pg_collection is None: # Use parallel_state groups dp_group = parallel_state.get_data_parallel_group( - with_context_parallel=False, partial_data_parallel=False + with_context_parallel=False, with_gtp_remat=False, partial_data_parallel=False ) dp_cp_group = parallel_state.get_data_parallel_group( - with_context_parallel=True, partial_data_parallel=False + with_context_parallel=True, with_gtp_remat=False, partial_data_parallel=False ) intra_dp_cp_group = parallel_state.get_data_parallel_group( - with_context_parallel=True, partial_data_parallel=True + with_context_parallel=True, with_gtp_remat=False, partial_data_parallel=True ) - expt_dp_group = parallel_state.get_expert_data_parallel_group() + expt_dp_group = parallel_state.get_expert_data_parallel_group(with_gtp_remat=False) intra_expt_dp_group = parallel_state.get_expert_data_parallel_group( - partial_expert_data_parallel=True + with_gtp_remat=False, partial_expert_data_parallel=True + ) + gtp_remat_group = parallel_state.get_gtp_weight_remat_group(check_initialized=False) + expt_gtp_remat_group = parallel_state.get_expert_gtp_weight_remat_group( + check_initialized=False ) intra_dist_opt_group = parallel_state.get_intra_distributed_optimizer_instance_group() - # Gloo groups - if use_gloo_process_groups: + # Gloo is not built under GTP_remat (the GTP_remat optimizer uses DCP); fetching the + # absent group would assert, so gate on gtp_active and leave the Gloo groups None. + gtp_active = (gtp_remat_group is not None and gtp_remat_group.size() > 1) or ( + expt_gtp_remat_group is not None and expt_gtp_remat_group.size() > 1 + ) + if use_gloo_process_groups and not gtp_active: intra_dp_cp_group_gloo = parallel_state.get_data_parallel_group_gloo( with_context_parallel=True, partial_data_parallel=True ) @@ -323,6 +405,9 @@ def setup_process_groups_for_optimizer( # Model communication groups mp_group = parallel_state.get_model_parallel_group() expt_tp_pp_group = parallel_state.get_expert_tensor_model_pipeline_parallel_group() + expt_tp_pp_with_egtp_remat_group = ( + parallel_state.get_expert_tensor_model_pipeline_parallel_group(with_egtp_remat=True) + ) # Inter distributed optimizer group if hasattr(model_chunks[0], 'ddp_config'): @@ -338,14 +423,15 @@ def setup_process_groups_for_optimizer( else: # Use provided process group collection with validation and fallbacks + pg_set = vars(pg_collection) # 1. dp group - this is always required - if not hasattr(pg_collection, 'dp'): + if 'dp' not in pg_set: raise ValueError("dp process group is required but not provided in pg_collection") dp_group = pg_collection.dp # 2. dp_cp group: fallback logic based on context_parallel_size - if hasattr(pg_collection, 'dp_cp'): + if 'dp_cp' in pg_set: dp_cp_group = pg_collection.dp_cp else: model_config = get_model_config(model_chunks[0]) @@ -360,7 +446,7 @@ def setup_process_groups_for_optimizer( ) # 3. Handle expert data parallel group - if not hasattr(pg_collection, 'expt_dp'): + if 'expt_dp' not in pg_set: raise ValueError( "expt_dp process group is required but not provided in pg_collection. " "Please explicitly set it to None if you don't need it." @@ -381,10 +467,10 @@ def setup_process_groups_for_optimizer( else: # With multiple optimizer instances, both groups must be provided if not ( - hasattr(pg_collection, 'intra_dp_cp') - and hasattr(pg_collection, 'intra_expt_dp') - and hasattr(pg_collection, 'inter_dist_opt') - and hasattr(pg_collection, 'intra_dist_opt') + 'intra_dp_cp' in pg_set + and 'intra_expt_dp' in pg_set + and 'inter_dist_opt' in pg_set + and 'intra_dist_opt' in pg_set ): raise ValueError( "intra_dp_cp, intra_expt_dp, inter_dist_opt, and intra_dist_opt " @@ -396,7 +482,7 @@ def setup_process_groups_for_optimizer( inter_dist_opt_group = pg_collection.inter_dist_opt if ddp_config.use_distributed_optimizer: - if not hasattr(pg_collection, 'intra_dist_opt'): + if 'intra_dist_opt' not in pg_set: raise ValueError( "intra_dist_opt process group is required but not provided in " "pg_collection. Please explicitly set it to None if you don't need it." @@ -412,7 +498,7 @@ def setup_process_groups_for_optimizer( intra_dist_opt_group = None # 5. Model communication groups - if not hasattr(pg_collection, 'mp'): + if 'mp' not in pg_set: raise ValueError( "mp process group is required but not provided in pg_collection. " "Please explicitly set it to None if you don't need it." @@ -420,13 +506,25 @@ def setup_process_groups_for_optimizer( mp_group = pg_collection.mp # Expert tensor-model-pipeline group for MoE - if not hasattr(pg_collection, 'tp_ep_pp'): + if 'tp_ep_pp' not in pg_set: raise ValueError( "tp_ep_pp process group is required but not provided in pg_collection. " "Please explicitly set it to None if you don't need it." ) expt_tp_pp_group = pg_collection.tp_ep_pp + # EGTP-MERGED variant of tp_ep_pp: includes the egtp axis, so each EGTP peer gets a + # distinct rank — used for the distopt ShardedObject keys. Falls back to tp_ep_pp + # when not provided. + if 'tp_ep_pp_with_egtp_remat' in pg_set: + expt_tp_pp_with_egtp_remat_group = pg_collection.tp_ep_pp_with_egtp_remat + else: + expt_tp_pp_with_egtp_remat_group = expt_tp_pp_group + + # GTP weight-shard groups (None when inactive); used to detect whether GTP is on. + gtp_remat_group = getattr(pg_collection, 'gtp_remat', None) + expt_gtp_remat_group = getattr(pg_collection, 'expt_gtp_remat', None) + # Gloo groups - not supported when pg_collection is provided if use_gloo_process_groups: raise ValueError( @@ -442,8 +540,11 @@ def setup_process_groups_for_optimizer( 'intra_dp_cp_group': intra_dp_cp_group, 'expt_dp_group': expt_dp_group, 'intra_expt_dp_group': intra_expt_dp_group, + 'gtp_remat_group': gtp_remat_group, + 'expt_gtp_remat_group': expt_gtp_remat_group, 'mp_group': mp_group, 'expt_tp_pp_group': expt_tp_pp_group, + 'expt_tp_pp_with_egtp_remat_group': expt_tp_pp_with_egtp_remat_group, 'inter_dist_opt_group': inter_dist_opt_group, 'intra_dist_opt_group': intra_dist_opt_group, 'intra_dp_cp_group_gloo': intra_dp_cp_group_gloo, @@ -478,21 +579,29 @@ def setup_process_groups_for_ddp( # Use parallel_state groups return { 'dp_group': parallel_state.get_data_parallel_group( - with_context_parallel=False, partial_data_parallel=False + with_context_parallel=False, with_gtp_remat=False, partial_data_parallel=False ), 'dp_cp_group': parallel_state.get_data_parallel_group( - with_context_parallel=True, partial_data_parallel=False + with_context_parallel=True, with_gtp_remat=False, partial_data_parallel=False ), 'intra_dp_cp_group': parallel_state.get_data_parallel_group( - with_context_parallel=True, partial_data_parallel=True + with_context_parallel=True, with_gtp_remat=False, partial_data_parallel=True + ), + 'expt_dp_group': parallel_state.get_expert_data_parallel_group( + with_gtp_remat=False ), - 'expt_dp_group': parallel_state.get_expert_data_parallel_group(), 'intra_expt_dp_group': parallel_state.get_expert_data_parallel_group( - partial_expert_data_parallel=True + with_gtp_remat=False, partial_expert_data_parallel=True ), 'tp_group': parallel_state.get_tensor_model_parallel_group(), + 'gtp_remat_group': parallel_state.get_gtp_weight_remat_group( + check_initialized=False + ), 'pp_group': parallel_state.get_pipeline_model_parallel_group(), 'ep_group': parallel_state.get_expert_model_parallel_group(), + 'expt_gtp_remat_group': parallel_state.get_expert_gtp_weight_remat_group( + check_initialized=False + ), 'inter_dist_opt_group': ( parallel_state.get_inter_distributed_optimizer_instance_group() if ddp_config.num_distributed_optimizer_instances > 1 @@ -507,14 +616,15 @@ def setup_process_groups_for_ddp( else: # Use provided process group collection with validation and fallbacks result = {} + pg_set = vars(pg_collection) # 1. dp group - this is always required - if not hasattr(pg_collection, 'dp'): + if 'dp' not in pg_set: raise ValueError("dp process group is required but not provided in pg_collection") result['dp_group'] = pg_collection.dp # 2. dp_cp group: fallback logic based on context_parallel_size - if hasattr(pg_collection, 'dp_cp'): + if 'dp_cp' in pg_set: result['dp_cp_group'] = pg_collection.dp_cp else: cp_size = getattr(config, 'context_parallel_size', 1) @@ -550,9 +660,9 @@ def setup_process_groups_for_ddp( else: # With multiple optimizer instances, groups must be provided if not ( - hasattr(pg_collection, 'intra_dp_cp') - and hasattr(pg_collection, 'intra_expt_dp') - and hasattr(pg_collection, 'inter_dist_opt') + 'intra_dp_cp' in pg_set + and 'intra_expt_dp' in pg_set + and 'inter_dist_opt' in pg_set ): raise ValueError( "intra_dp_cp, intra_expt_dp, and inter_dist_opt " @@ -564,13 +674,7 @@ def setup_process_groups_for_ddp( result['inter_dist_opt_group'] = pg_collection.inter_dist_opt # 5. Model parallel groups (DDP-specific: tp, pp, ep instead of mp, expt_tp_pp) - if not all( - [ - hasattr(pg_collection, 'tp'), - hasattr(pg_collection, 'pp'), - hasattr(pg_collection, 'ep'), - ] - ): + if not all(['tp' in pg_set, 'pp' in pg_set, 'ep' in pg_set]): raise ValueError( "tp, pp and ep process groups are required but not provided in pg_collection" ) @@ -578,6 +682,10 @@ def setup_process_groups_for_ddp( result['pp_group'] = pg_collection.pp result['ep_group'] = pg_collection.ep + # GTP weight-shard groups (None when inactive); used to detect whether GTP is on. + result['gtp_remat_group'] = getattr(pg_collection, 'gtp_remat', None) + result['expt_gtp_remat_group'] = getattr(pg_collection, 'expt_gtp_remat', None) + return result diff --git a/megatron/core/recompute.py b/megatron/core/recompute.py index 8974efc1311..4af57efaf4e 100644 --- a/megatron/core/recompute.py +++ b/megatron/core/recompute.py @@ -12,10 +12,10 @@ from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.transformer_layer import TransformerLayer -te_checkpoint = None - if HAVE_TE: from megatron.core.extensions.transformer_engine import te_checkpoint +else: + te_checkpoint = None def checkpointed_forward( @@ -51,10 +51,27 @@ def checkpointed_forward( extract_layer_indices = set() intermediate_hidden_states: List[Tensor] = [] + # Wrap non-dual RoPE to tuple to unify custom_forward interface. + is_dual_rope = isinstance(rotary_pos_emb, (tuple, list)) + assert not is_dual_rope or len(rotary_pos_emb) == 2, "Dual RoPE input length is not equal to 2" + rotary_pos_emb = rotary_pos_emb if is_dual_rope else (None, rotary_pos_emb) + def custom(start: int, end: int): def custom_forward( - hidden_states, attention_mask, context, context_mask, rotary_pos_emb, padding_mask=None + hidden_states, + attention_mask, + context, + context_mask, + rotary_pos_emb_local, + rotary_pos_emb_global, + padding_mask=None, ): + rotary_pos_emb = ( + (rotary_pos_emb_local, rotary_pos_emb_global) + if is_dual_rope + else rotary_pos_emb_global + ) + for index in range(start, end): # Use self.layers[index] (not self._get_layer) so this # function works for both TransformerBlock and HybridStack. @@ -116,7 +133,8 @@ def custom_forward( def chunk_runner(start: int, end: int, use_checkpoint: bool): nonlocal hidden_states, context cf = custom(start, end) - args = (hidden_states, attention_mask, context, context_mask, rotary_pos_emb, padding_mask) + # Unpack the RoPE tuple as torch cannot save tuples for backward pass. + args = (hidden_states, attention_mask, context, context_mask, *rotary_pos_emb, padding_mask) if use_checkpoint: # Precision-aware activation checkpoint: TE under FP8/FP4, # tensor_parallel under BF16/FP16/FP32. diff --git a/megatron/core/resharding/README.md b/megatron/core/resharding/README.md index 875d82d2c95..134e9c9225c 100644 --- a/megatron/core/resharding/README.md +++ b/megatron/core/resharding/README.md @@ -10,7 +10,8 @@ inference model that may use a different parallelism layout. ``` refit.py High-level API: swap_model_weights, caching, MXFP8 auto-detection | -planner.py Centralized plan builder (rank 0 builds, scatters to all) +planner.py Local plan builder (every rank all-gathers metadata, replays + the same deterministic schedule, keeps only its own ops) | execution.py Submits send/recv ops to a CopyService, handles writebacks | @@ -79,6 +80,7 @@ swap_model_weights(None, None, "nccl", | `nccl` | GPU P2P via `batch_isend_irecv` | Intra-node / single cluster | Lowest latency; default choice | | `gloo` | CPU-staged via Gloo PG | Cross-cluster / multi-node | Higher latency; works where NCCL cross-cluster doesn't | | `nvshmem` | Pipelined NVSHMEM puts | High-throughput intra-node | Requires NVSHMEM; uses double-buffered kernel pipeline | +| `nixl` | GPU RDMA via NIXL (UCX), sender-initiated WRITE | Cross-cluster / non-collocated | Requires NIXL; transfers GPU memory directly (no host staging) | All backends detect same-rank (local) transfers via `task_id` and short-circuit them into direct `tensor.copy_()` instead of going @@ -87,15 +89,26 @@ through the network stack. ## How the Reshard Plan Works 1. Each rank extracts parameter metadata (shape, sharding, TP/EP/PP groups). -2. Metadata is gathered to rank 0 via `dist.gather_object()`. -3. Rank 0 builds a complete transfer schedule: - - For each destination param, finds the matching source param(s) by name. - - Routes to a dimension-specific planner (LCM tiling for standard TP, +2. Metadata is all-gathered so **every** rank has the full picture + (`dist.all_gather_object()`) — no rank-0 bottleneck, no scatter. +3. Every rank independently replays the **same deterministic schedule** + (`_iter_global_transfer_ops`): + - Iterate destination ranks, then each rank's destination params in gathered + order; for each destination param, find the matching source param(s) by name. + - Route to a dimension-specific planner (LCM tiling for standard TP, block-interleaved for partitioned params like Mamba `in_proj`). - - Produces `TransferOp` pairs with globally unique `task_id` values. -4. Plans are scattered back; each rank receives only its own send/recv ops. + - Assign a monotonic `task_id` per sub-op. Because the iteration order and + counter are a pure function of the gathered metadata, the send op computed + on the sender and the recv op computed on the receiver get the **same** + `task_id` without any central authority. +4. Each rank keeps only the ops where it is the sender or receiver. 5. The plan is cached so repeated refits skip steps 1-4. +The deterministic schedule stays stable when a larger roster is supplied: existing +transfers keep their `task_id`s and newly appended destination ranks receive new +ones. Live process-group membership changes and their orchestration remain future +work; this module does not currently add or remove ranks from a running group. + ## MXFP8 Transform When the target model uses `transformer_impl='inference_optimized'` with @@ -144,11 +157,12 @@ attribute with the following groups: | File | Role | |------|------| | `refit.py` | Public API, caching, MXFP8 auto-detection | -| `planner.py` | Centralized plan builder (metadata, LCM/block-interleaved planners) | +| `planner.py` | Local deterministic plan builder (metadata, LCM/block-interleaved planners) | | `execution.py` | Plan executor (send/recv submission, writeback, format conversion) | | `transforms.py` | `ReshardTransform` base class, `MXFP8ReshardTransform` | | `utils.py` | `TransferOp`, `ReshardPlan`, `ParameterMetadata`, `ShardingDescriptor` | | `copy_services/nccl_copy_service.py` | NCCL backend | | `copy_services/gloo_copy_service.py` | Gloo backend | +| `copy_services/nixl_copy_service.py` | NIXL/UCX backend | | `copy_services/nvshmem_copy_service.py` | NVSHMEM backend (delegates to `nvshmem_copy_service/`) | | `nvshmem_copy_service/` | Full NVSHMEM implementation (planning, memory, kernels, pipeline) | diff --git a/megatron/core/resharding/__init__.py b/megatron/core/resharding/__init__.py index 8c59b6ef809..7fc122c2400 100644 --- a/megatron/core/resharding/__init__.py +++ b/megatron/core/resharding/__init__.py @@ -1,6 +1,11 @@ # Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. from .execution import execute_reshard_plan -from .planner import build_centralized_reshard_plan +from .planner import ( + build_centralized_reshard_plan, + build_local_reshard_plan, + build_plan_from_rosters, + index_metadata_rosters, +) from .refit import ( clear_service_cache, get_or_create_service, @@ -12,6 +17,9 @@ __all__ = [ "build_centralized_reshard_plan", + "build_local_reshard_plan", + "build_plan_from_rosters", + "index_metadata_rosters", "execute_reshard_plan", "MXFP8ReshardTransform", "ReshardTransform", diff --git a/megatron/core/resharding/copy_services/__init__.py b/megatron/core/resharding/copy_services/__init__.py index c85ebb0ece2..18ee6afa2fe 100644 --- a/megatron/core/resharding/copy_services/__init__.py +++ b/megatron/core/resharding/copy_services/__init__.py @@ -4,6 +4,13 @@ from .base import CopyService from .gloo_copy_service import GlooCopyService from .nccl_copy_service import NCCLCopyService +from .nixl_copy_service import NixlCopyService from .nvshmem_copy_service import NVSHMEMCopyService -__all__ = ["CopyService", "GlooCopyService", "NCCLCopyService", "NVSHMEMCopyService"] +__all__ = [ + "CopyService", + "GlooCopyService", + "NCCLCopyService", + "NixlCopyService", + "NVSHMEMCopyService", +] diff --git a/megatron/core/resharding/copy_services/base.py b/megatron/core/resharding/copy_services/base.py index 00dc884767c..98404d28015 100644 --- a/megatron/core/resharding/copy_services/base.py +++ b/megatron/core/resharding/copy_services/base.py @@ -37,6 +37,11 @@ class CopyService(ABC): remote transfers simply ignore it. """ + # Most torch.distributed backends retain the executor's historical + # process-group rendezvous after run(). Backends whose run() protocol already + # establishes completion across every participating peer can opt out. + requires_process_group_barrier = True + def __init__(self, group=None): self.group = group # group.rank()/size() supports cross-cluster ProcessGroups where members diff --git a/megatron/core/resharding/copy_services/nixl_copy_service.py b/megatron/core/resharding/copy_services/nixl_copy_service.py new file mode 100644 index 00000000000..917ad6fbf33 --- /dev/null +++ b/megatron/core/resharding/copy_services/nixl_copy_service.py @@ -0,0 +1,277 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +from __future__ import annotations + +import logging +import os +import time +from typing import Dict, List, Optional, Tuple + +import torch +import torch.distributed as dist + +from .base import CopyService, RecvOp, SendOp, match_local_ops_by_task_id + +logger = logging.getLogger(__name__) + +# Transports that let UCX read/write CUDA memory. Without one of these UCX sees +# a GPU pointer as host memory and segfaults mid-transfer. +_CUDA_UCX_TRANSPORTS = ("cuda_copy", "cuda_ipc") + +# (addr, len_bytes, device_id) for a registered region. Exchanged between ranks +# so a sender can WRITE straight into a receiver's destination buffer. +_MemDesc = Tuple[int, int, int] + + +def _ensure_cuda_ucx_transports() -> None: + """Add cuda transports to UCX_TLS if it's pinned to a host-only allowlist. + + UCX reads UCX_TLS once at agent init. Deployments often set it to e.g. "tcp", + which can't touch GPU memory. Only augment a plain inclusion list; leave an + unset value (defaults to "all") or an exclusion list ("^...") alone. + """ + tls = os.environ.get("UCX_TLS") + if not tls: + return + tokens = [t.strip() for t in tls.split(",") if t.strip()] + if not tokens or tokens[0].startswith("^") or any("cuda" in t for t in tokens): + return + os.environ["UCX_TLS"] = tls + "," + ",".join(_CUDA_UCX_TRANSPORTS) + logger.warning("UCX_TLS=%r has no cuda transport; using %r", tls, os.environ["UCX_TLS"]) + + +class NixlCopyService(CopyService): + """Refit transport over NIXL (UCX/RDMA), for cross-cluster non-collocated refit. + + Each rank runs a NIXL agent. To WRITE into a peer it needs that peer's agent + metadata (connection info) plus a {task_id: (addr, len, dev)} map of the peer's + registered recv buffers. A one-time torch all-gather builds and caches this + peer table. Buffers are registered locally, not exchanged. + + Every refit after that is pure NIXL and sender-driven: a receiver signals each + source that its buffers are free ("ready"); the source waits, syncs its weights, + and issues one WRITE per receiver, each carrying a "data" notification. Those two + notifications order producer and consumer per refit, so there's no barrier and no + per-refit collective. Notifications are tagged with a per-refit sequence, so stale + ones are ignored. Same-rank transfers skip NIXL and copy directly. + + Registered buffers are assumed address-stable across refits; if a recv address + changes after setup, call clear_service_cache() to rebuild. + """ + + requires_process_group_barrier = False + + def __init__(self, group=None, agent_name: Optional[str] = None): + super().__init__(group=group) + # Object collectives (the one-time handshake) run on this group to support + # cross-world PGs. + self.group = group + + try: + from nixl._api import nixl_agent, nixl_agent_config + except ImportError as e: + raise ImportError( + "NixlCopyService requires the 'nixl' package; install it or use " + "another refit backend (nccl/gloo/nvshmem)." + ) from e + + _ensure_cuda_ucx_transports() + + # Name by group rank, which is unique across the (possibly cross-world) + # group. dist.get_rank() would collide: separate worlds each have a rank 0. + self.agent_name = agent_name or f"refit-nixl-rank-{self.rank}" + # This backend is intentionally UCX-only: refit moves CPU/CUDA tensors + # directly between ranks and does not use NIXL's storage-oriented plugins. + self.agent = nixl_agent(self.agent_name, nixl_agent_config(backends=["UCX"])) + + self.send_ops: List[SendOp] = [] + self.recv_ops: List[RecvOp] = [] + self._copy_stream = torch.cuda.Stream() + + # Peer table {group rank -> (agent_name, agent_metadata, recv_descs)}, + # populated by the first run's handshake. + self._remote_agent_names: Dict[int, str] = {} # group rank -> agent name + self._gathered: Optional[Dict[int, tuple]] = None + self._recv_descs: Optional[Dict[int, _MemDesc]] = None + # RDMA registrations, kept across refits; send and recv tracked separately + # so a changing side doesn't re-pin the other. + self._reg: Dict[str, tuple] = {} # 'send'/'recv' -> (handle, signature) + # Notifications are tagged (kind, seq). _future_notifs buffers tags for + # later refits that we're not draining yet. + self._seq = 0 + self._future_notifs: Dict[Tuple[str, int], int] = {} + + def submit_send(self, src_tensor: torch.Tensor, dest_rank: int, task_id: Optional[int] = None): + self.send_ops.append(SendOp(task_id=task_id, tensor=src_tensor, dest_rank=dest_rank)) + + def submit_recv(self, dest_tensor: torch.Tensor, src_rank: int, task_id: Optional[int] = None): + self.recv_ops.append(RecvOp(task_id=task_id, tensor=dest_tensor, src_rank=src_rank)) + + @staticmethod + def _mem_desc(tensor: torch.Tensor) -> _MemDesc: + if not tensor.is_contiguous(): + raise RuntimeError("NixlCopyService requires contiguous tensors") + dev = tensor.get_device() # -1 for host tensors; NIXL addresses DRAM as device 0 + return (tensor.data_ptr(), tensor.numel() * tensor.element_size(), dev if dev >= 0 else 0) + + def _register(self, which: str, tensors: List[torch.Tensor]) -> None: + # Re-register only when the set of regions changes. + sig = tuple((t.data_ptr(), t.numel() * t.element_size()) for t in tensors) + cached = self._reg.get(which) + if cached is not None and cached[1] == sig: + return + if cached is not None and cached[0] is not None: + self.agent.deregister_memory(cached[0]) + handle = self.agent.register_memory(tensors) if tensors else None + self._reg[which] = (handle, sig) + + def _handshake(self, recv_descs: Dict[int, _MemDesc]) -> None: + # Initial bootstrap over the torch group: share agent metadata and recv + # descriptors, connect to every peer, and cache the peer table. This is + # NIXL's one torch collective. + payload = (self.agent_name, self.agent.get_agent_metadata(), recv_descs) + gathered: List[Optional[tuple]] = [None] * self.world_size + dist.all_gather_object(gathered, payload, group=self.group) + self._gathered = {rank: entry for rank, entry in enumerate(gathered) if entry is not None} + for rank, (name, metadata, _descs) in self._gathered.items(): + if rank != self.rank: + self._remote_agent_names[rank] = self.agent.add_remote_agent(metadata) or name + self._recv_descs = recv_descs + + def _do_local_copies(self) -> None: + # Collocated (same-rank) transfers never hit the network. + local_sends = [op for op in self.send_ops if op.dest_rank == self.rank] + local_recvs = [op for op in self.recv_ops if op.src_rank == self.rank] + if not local_sends and not local_recvs: + return + pairs = match_local_ops_by_task_id(local_sends, local_recvs, "NixlCopyService", self.rank) + with torch.no_grad(), torch.cuda.stream(self._copy_stream): + for send_op, recv_op in pairs: + recv_op.tensor.copy_(send_op.tensor) + + def _plan_writes(self, remote_sends: List[SendOp]): + # Group writes by destination agent: one WRITE per receiver pushes all its + # regions at once, with local[i] landing in remote[i]. + if self._gathered is None: + raise RuntimeError("NixlCopyService: handshake has not completed") + by_dst: Dict[int, Tuple[List[torch.Tensor], List[_MemDesc]]] = {} + for op in remote_sends: + dst_entry = self._gathered.get(op.dest_rank) + if dst_entry is None: + raise RuntimeError(f"NixlCopyService: no metadata from dst rank {op.dest_rank}") + remote_desc = dst_entry[2].get(op.task_id) + if remote_desc is None: + raise RuntimeError( + f"NixlCopyService: dst rank {op.dest_rank} missing task_id {op.task_id}" + ) + local_list, remote_list = by_dst.setdefault(op.dest_rank, ([], [])) + local_list.append(op.tensor) + remote_list.append(remote_desc) + return by_dst + + @staticmethod + def _notif(kind: str, seq: int) -> bytes: + return f"{kind}{seq}".encode() + + @staticmethod + def _parse_notif(m: bytes) -> Tuple[str, int]: + return chr(m[0]), int(m[1:]) + + def _await_notifs(self, kind: str, expected: int, seq: int) -> None: + # Wait for `expected` notifications of this kind ('R' ready / 'D' data) + # tagged with this refit. Buffer any tagged otherwise — e.g. a data notif + # arriving while we're still collecting ready signals. + if expected == 0: + return + want = (kind, seq) + got = self._future_notifs.pop(want, 0) + while got < expected: + for _agent, msgs in self.agent.get_new_notifs().items(): + for m in msgs: + tag = self._parse_notif(m) + if tag == want: + got += 1 + else: + self._future_notifs[tag] = self._future_notifs.get(tag, 0) + 1 + if got < expected: + time.sleep(0) # NIXL delivers notifs on its own thread + + def run(self): + remote_sends = [op for op in self.send_ops if op.dest_rank != self.rank] + remote_recvs = [op for op in self.recv_ops if op.src_rank != self.rank] + seq = self._seq + + # Overlaps with the writes below. + self._do_local_copies() + + # Both sides of a WRITE must be registered: our send tensors (we read them) + # and our recv buffers (peers write into them). + self._register("send", [op.tensor for op in remote_sends]) + self._register("recv", [op.tensor for op in remote_recvs]) + + recv_descs = {op.task_id: self._mem_desc(op.tensor) for op in remote_recvs} + if self._gathered is None: + self._handshake(recv_descs) + elif recv_descs != self._recv_descs: + raise RuntimeError( + "NixlCopyService: recv tensor addresses changed after setup; " + "call clear_service_cache() to rebuild the handshake." + ) + + # Tell each source our destination buffers are free for this refit; this is + # what orders a source's write after we've consumed the previous one. + ready = self._notif("R", seq) + for src_rank in {op.src_rank for op in remote_recvs}: + self.agent.send_notif(self._remote_agent_names[src_rank], ready) + + # As a source: wait for every receiver's ready, then push once the weights + # are finished (a peer must not read a half-written buffer). + self._await_notifs("R", len({op.dest_rank for op in remote_sends}), seq) + if remote_sends: + torch.cuda.current_stream().synchronize() + + data = self._notif("D", seq) + handles = [] + for dst_rank, (local_tensors, remote_descs) in self._plan_writes(remote_sends).items(): + # local and remote are the same memory class (GPU->GPU for refit); the + # remote tuples carry no type, so take it from the send tensor. + mem_type = "cuda" if local_tensors[0].is_cuda else "cpu" + local_xfer = self.agent.get_xfer_descs(local_tensors) + remote_xfer = self.agent.get_xfer_descs(remote_descs, mem_type=mem_type) + handle = self.agent.initialize_xfer( + "WRITE", local_xfer, remote_xfer, self._remote_agent_names[dst_rank], notif_msg=data + ) + if self.agent.transfer(handle) == "ERR": + raise RuntimeError(f"NixlCopyService: WRITE to dst {dst_rank} failed to start") + handles.append((dst_rank, handle)) + + # NIXL's asynchronous Python API exposes transfer completion through + # check_xfer_state(); the official examples poll it until DONE. Transfers + # progress in the backend, so yield the Python thread between checks. + for dst_rank, handle in handles: + state = self.agent.check_xfer_state(handle) + while state not in ("DONE", "ERR"): + time.sleep(0) + state = self.agent.check_xfer_state(handle) + if state == "ERR": + raise RuntimeError(f"NixlCopyService: WRITE to dst {dst_rank} errored") + self.agent.release_xfer_handle(handle) + self._await_notifs("D", len({op.src_rank for op in remote_recvs}), seq) + + torch.cuda.current_stream().wait_stream(self._copy_stream) + self._seq += 1 + self.send_ops.clear() + self.recv_ops.clear() + + def close(self) -> None: + for handle, _sig in self._reg.values(): + if handle is not None: + self.agent.deregister_memory(handle) + self._reg.clear() + for name in self._remote_agent_names.values(): + try: + self.agent.remove_remote_agent(name) + except Exception: + pass + self._remote_agent_names.clear() + self._gathered = None + self._recv_descs = None diff --git a/megatron/core/resharding/execution.py b/megatron/core/resharding/execution.py index 8b06ef1cd14..ce45f124b58 100644 --- a/megatron/core/resharding/execution.py +++ b/megatron/core/resharding/execution.py @@ -64,7 +64,7 @@ def execute_reshard_plan( transform: Optional[ReshardTransform] = None, ) -> None: """ - Execute a reshard plan (from centralized controller). + Execute a reshard plan (built locally on each rank). A communication service must be provided to abstract transport. Expected service API: submit_send(tensor, dest_rank, task_id), submit_recv(tensor, src_rank, task_id), run(). @@ -195,7 +195,8 @@ def get_sendable(param_name: str, param: torch.nn.Parameter) -> torch.Tensor: logger.info(f"Executing {len(plan.send_ops)} sends + {len(plan.recv_ops)} recvs") service.run() torch.cuda.synchronize() - dist.barrier(group=group) + if service.requires_process_group_barrier: + dist.barrier(group=group) # Write back received buffers into their destination parameter slices. # diff --git a/megatron/core/resharding/nvshmem_copy_service/core/gpu_resource_manager.py b/megatron/core/resharding/nvshmem_copy_service/core/gpu_resource_manager.py index 0a8c94afe65..f10bcfd5680 100644 --- a/megatron/core/resharding/nvshmem_copy_service/core/gpu_resource_manager.py +++ b/megatron/core/resharding/nvshmem_copy_service/core/gpu_resource_manager.py @@ -109,17 +109,18 @@ def init(self, group=None) -> None: "Could not determine nvidia-nvshmem-cu12 package version for NVSHMEM safety check." ) - # Recommend a conservative CTA limit for stability when team counts grow. + # This path can hang during initialization when the CTA limit is higher. max_ctas = os.environ.get("NVSHMEM_MAX_CTAS") if max_ctas != "2": - logger.warning( - "Recommended NVSHMEM_MAX_CTAS=2 for this path. Current value is %r.", max_ctas + raise RuntimeError( + "NVSHMEM_MAX_CTAS must be set to '2' for the NVSHMEM copy service; " + f"got {max_ctas!r}." ) # torch.distributed must be initialized before calling this if not dist.is_initialized(): raise RuntimeError( - "torch.distributed must be initialized before " "GPUResourceManager.init()" + "torch.distributed must be initialized before GPUResourceManager.init()" ) # Get current CUDA device (already set by caller based on LOCAL_RANK) diff --git a/megatron/core/resharding/nvshmem_copy_service/core/pipeline_executor.py b/megatron/core/resharding/nvshmem_copy_service/core/pipeline_executor.py index b4e183b886c..7bee9724413 100644 --- a/megatron/core/resharding/nvshmem_copy_service/core/pipeline_executor.py +++ b/megatron/core/resharding/nvshmem_copy_service/core/pipeline_executor.py @@ -7,6 +7,7 @@ and proper stream synchronization. """ +import time from typing import Dict, List, Optional from ..compat import ensure_nvshmem_compat @@ -117,6 +118,17 @@ def execute_pipeline( """ PELogger.info(f"Executing pipeline: {num_iterations} iterations") + pipeline_start = time.perf_counter() + wait_unpack_seconds = 0.0 + put_submit_seconds = 0.0 + quiet_submit_seconds = 0.0 + barrier_submit_seconds = 0.0 + barrier_host_sync_seconds = 0.0 + wait_pack_seconds = 0.0 + final_unpack_seconds = 0.0 + slowest_barrier_sync_iteration = -1 + slowest_barrier_sync_seconds = 0.0 + # Priming: Pack iteration 0 (async, no CPU sync needed — # step 3 uses GPU-level event wait for pack→put ordering) if num_iterations > 0 and iter_schedules[0]["send"]: @@ -156,7 +168,7 @@ def execute_pipeline( next_batch = iter_schedules[i + 1]["send"] assert next_batch is not None PELogger.debug( - f" Pack next (iter {i+1}): {len(next_batch.tasks)} tasks " + f" Pack next (iter {i + 1}): {len(next_batch.tasks)} tasks " f"→ PE {next_batch.dest_pe}" ) self._launch_pack(i + 1, next_batch) @@ -168,7 +180,7 @@ def execute_pipeline( prior_batch = iter_schedules[i - 1]["recv"] assert prior_batch is not None PELogger.debug( - f" Unpack prior (iter {i-1}): {prior_batch.total_size} bytes " + f" Unpack prior (iter {i - 1}): {prior_batch.total_size} bytes " f"← PE {prior_batch.src_pe}" ) # GPU-level event wait: ensures send_stream's barrier_all from @@ -191,39 +203,50 @@ def execute_pipeline( # the NIC's DMA engine. self.torch_send_stream_wrapper.wait_event(self.pack_events[slot]) + put_start = time.perf_counter() nvshmem.core.put( self.buffer_manager.recv_slots[slot][0:transfer_size], self.buffer_manager.send_slots[slot][0:transfer_size], batch.dest_pe, stream=self.send_stream, ) + put_submit_seconds += time.perf_counter() - put_start nvtx_range_pop("Step 3: Send Current") # Step 4a: Wait for prior unpack to complete BEFORE the barrier. nvtx_range_push("Step 4a: Wait Unpack") if has_prior_recv: + wait_start = time.perf_counter() self.unpack_events[(i - 1) % 2].synchronize() + wait_unpack_seconds += time.perf_counter() - wait_start nvtx_range_pop("Step 4a: Wait Unpack") # Ensure all NVSHMEM operations on send_stream complete (stream-ordered) + quiet_start = time.perf_counter() nvshmem.core.quiet(stream=self.send_stream) + quiet_submit_seconds += time.perf_counter() - quiet_start - # Step 4b: Global barrier + CPU sync + record event + # Step 4b: Global barrier + CPU sync + record event. nvtx_range_push("Step 4b: Barrier") + barrier_start = time.perf_counter() nvshmem.core.barrier_all(stream=self.send_stream) - # CPU-sync the send_stream to ensure barrier_all has actually - # completed (not just submitted). Without this, the barrier_event - # can fire before RDMA data from the remote PE is visible, because - # stream-ordered operations are only guaranteed to be submitted, - # not completed, when the event is recorded. + barrier_submit_seconds += time.perf_counter() - barrier_start + sync_start = time.perf_counter() self.torch_send_stream_wrapper.synchronize() + sync_seconds = time.perf_counter() - sync_start + barrier_host_sync_seconds += sync_seconds + if sync_seconds > slowest_barrier_sync_seconds: + slowest_barrier_sync_iteration = i + slowest_barrier_sync_seconds = sync_seconds self.barrier_events[slot].record(stream=self.torch_send_stream_wrapper) nvtx_range_pop("Step 4b: Barrier") # Step 5: Wait for async pack to complete (double-buffer safety) nvtx_range_push("Step 5: Wait Pack") if has_next_send: + wait_start = time.perf_counter() self.pack_events[(i + 1) % 2].synchronize() + wait_pack_seconds += time.perf_counter() - wait_start nvtx_range_pop("Step 5: Wait Pack") nvtx_range_pop(nvtx_iter_msg) @@ -231,7 +254,7 @@ def execute_pipeline( # Final unpack for last iteration if num_iterations > 0 and iter_schedules[num_iterations - 1]["recv"]: nvtx_range_push("Final Unpack") - PELogger.debug(f"Final unpack: iteration {num_iterations-1}") + PELogger.debug(f"Final unpack: iteration {num_iterations - 1}") last_recv = iter_schedules[num_iterations - 1]["recv"] assert last_recv is not None # GPU-level event wait for NVSHMEM RDMA data visibility @@ -239,9 +262,25 @@ def execute_pipeline( self.barrier_events[(num_iterations - 1) % 2] ) self._launch_unpack(num_iterations - 1, last_recv) + wait_start = time.perf_counter() self.unpack_events[(num_iterations - 1) % 2].synchronize() + final_unpack_seconds += time.perf_counter() - wait_start nvtx_range_pop("Final Unpack") + pipeline_seconds = time.perf_counter() - pipeline_start + PELogger.info( + "Pipeline timing: " + f"total={pipeline_seconds:.6f}s " + f"wait_unpack={wait_unpack_seconds:.6f}s " + f"put_submit={put_submit_seconds:.6f}s " + f"quiet_submit={quiet_submit_seconds:.6f}s " + f"barrier_submit={barrier_submit_seconds:.6f}s " + f"barrier_host_sync={barrier_host_sync_seconds:.6f}s " + f"wait_pack={wait_pack_seconds:.6f}s " + f"final_unpack={final_unpack_seconds:.6f}s " + f"slowest_barrier_sync_iteration={slowest_barrier_sync_iteration} " + f"slowest_barrier_sync={slowest_barrier_sync_seconds:.6f}s" + ) PELogger.info(f"Pipeline complete: {num_iterations} iterations") def _launch_pack(self, iteration: int, batch: ScheduledBatch) -> None: diff --git a/megatron/core/resharding/nvshmem_copy_service/service.py b/megatron/core/resharding/nvshmem_copy_service/service.py index 332fe892c4e..4cf5151fc5f 100644 --- a/megatron/core/resharding/nvshmem_copy_service/service.py +++ b/megatron/core/resharding/nvshmem_copy_service/service.py @@ -8,6 +8,7 @@ GPU resource management, and pipelined execution. """ +import time from typing import Dict, List, Optional, Tuple from .compat import ensure_nvshmem_compat @@ -102,8 +103,12 @@ def init(self, log_level: str = "INFO") -> None: "nvshmem.core is not available. Please install nvshmem to use NVSHMEMCopyService." ) + init_start = time.perf_counter() + # Initialize GPU resources (NVSHMEM, device, streams) + phase_start = time.perf_counter() self.gpu_resources.init(group=self._group) + gpu_resources_seconds = time.perf_counter() - phase_start # Initialize logger after PE ID is known PELogger.init(self.my_pe, level=log_level) @@ -113,10 +118,13 @@ def init(self, log_level: str = "INFO") -> None: # buffer_manager.allocate() calls bytetensor() which is a collective operation # Without this barrier, early PEs call bytetensor() while late PEs # are still in init() -> deadlock + phase_start = time.perf_counter() nvshmem.core.barrier_all(stream=self.gpu_resources.send_stream) self.gpu_resources.send_stream.sync() # Ensure barrier completes on CPU + initial_barrier_seconds = time.perf_counter() - phase_start # Allocate double-buffered send/recv slots + phase_start = time.perf_counter() self.buffer_manager.allocate() # The .zero_() calls inside allocate() go to the default CUDA stream. # Sync it now so the zeros are fully committed before any NVShmem @@ -124,6 +132,7 @@ def init(self, log_level: str = "INFO") -> None: # Without this, a still-running zero() can race with the first # nvshmem.core.put() and overwrite received data. torch.cuda.synchronize() + buffer_allocation_seconds = time.perf_counter() - phase_start # Barrier to ensure all PEs complete buffer allocation before proceeding nvshmem.core.barrier_all(stream=self.gpu_resources.send_stream) @@ -131,7 +140,9 @@ def init(self, log_level: str = "INFO") -> None: PELogger.debug("Allocated double-buffered send/recv slots") # Load CUDA kernels + phase_start = time.perf_counter() self.kernel_launcher.load_kernels() + kernel_load_seconds = time.perf_counter() - phase_start PELogger.debug("Loaded CUDA kernels") # Cache CuPy stream wrappers for efficient kernel launching @@ -165,6 +176,14 @@ def init(self, log_level: str = "INFO") -> None: self.gpu_resources.unpack_stream.sync() self.gpu_resources.copy_stream.sync() + PELogger.info( + "Initialization timing: " + f"total={time.perf_counter() - init_start:.6f}s " + f"gpu_resources={gpu_resources_seconds:.6f}s " + f"initial_barrier={initial_barrier_seconds:.6f}s " + f"buffer_allocation={buffer_allocation_seconds:.6f}s " + f"kernel_load={kernel_load_seconds:.6f}s" + ) PELogger.info("Initialization complete") def register_send( @@ -224,6 +243,7 @@ def schedule(self) -> None: if not self.initialized: raise RuntimeError("RemoteCopyService not initialized") + schedule_start = time.perf_counter() PELogger.info( f"Starting schedule: {len(self.send_requests)} send requests, " f"{len(self.receive_requests)} receive requests" @@ -277,7 +297,10 @@ def schedule(self) -> None: ) self.pipeline_executor.set_events(self.pack_events, self.unpack_events, self.barrier_events) - PELogger.info(f"Schedule complete: {self.num_iterations} iterations ready") + PELogger.info( + f"Schedule complete: {self.num_iterations} iterations ready " + f"in {time.perf_counter() - schedule_start:.6f}s" + ) def run(self) -> None: """ @@ -295,6 +318,7 @@ def run(self) -> None: if self.iter_schedules is None: raise RuntimeError("Must call schedule() before run()") + run_start = time.perf_counter() PELogger.info(f"Starting execution: {self.num_iterations} iterations") # Start timing @@ -302,12 +326,16 @@ def run(self) -> None: # Global barrier before execution PELogger.debug("Barrier: Synchronizing all PEs before execution") + phase_start = time.perf_counter() nvshmem.core.barrier_all(stream=self.gpu_resources.send_stream) self.gpu_resources.send_stream.sync() + initial_barrier_seconds = time.perf_counter() - phase_start # Execute pipelined communication nvtx_range_push("execute_pipeline") + phase_start = time.perf_counter() self.pipeline_executor.execute_pipeline(self.iter_schedules, self.num_iterations) + pipeline_seconds = time.perf_counter() - phase_start nvtx_range_pop("execute_pipeline") # Global barrier after execution @@ -319,6 +347,12 @@ def run(self) -> None: # End timing range nvtx_range_pop("RemoteCopyService.run_total") + PELogger.info( + "Execution timing: " + f"total={time.perf_counter() - run_start:.6f}s " + f"initial_barrier={initial_barrier_seconds:.6f}s " + f"pipeline={pipeline_seconds:.6f}s" + ) def clear_requests(self) -> None: """ diff --git a/megatron/core/resharding/planner.py b/megatron/core/resharding/planner.py index 1eda91e914b..f89f39a4a0c 100644 --- a/megatron/core/resharding/planner.py +++ b/megatron/core/resharding/planner.py @@ -3,6 +3,7 @@ import logging import math +import warnings import torch import torch.distributed as dist @@ -301,170 +302,223 @@ def _determine_source_ranks_for_dst_param( return _finalize_dp_transfers(param_name, src_metadata, dst_metadata, my_global_rank) -def build_centralized_reshard_plan( +def _iter_global_transfer_ops( + dst_param_metadata_by_rank: dict[int, dict[str, ParameterMetadata]], + src_param_metadata: dict[str, list[ParameterMetadata]], +): + """Yield the whole reshard schedule in a deterministic order. + + The iteration order (dst rank ascending, then that rank's dst params in + gathered order, then per-source sub-ops) depends only on the rosters, so + replaying this on any rank produces the same sequence and assigns the same + task_id to the same transfer. That's what lets each rank build its own + send/recv ops while sender and receiver still agree on task_id. + + Ranks are taken from the roster keys rather than range(world_size), so a + sparse or growing rank set (nodes added later) rebuilds identically. + + Yields (task_id, dst_rank, src_rank, src_slice, dst_slice, src_metadata, + dst_metadata). PP is handled implicitly: each rank contributes metadata only + for the params it owns, and any source holding the same resolved_name can + serve as sender (with DP balancing). + """ + # Shared between a send and its recv; NVSHMEM builds a schedule from it and + # local copies match on it. + next_task_id = 0 + for dst_rank in sorted(dst_param_metadata_by_rank): + dst_rank_params = dst_param_metadata_by_rank[dst_rank] + for resolved_name, dst_metadata in dst_rank_params.items(): + src_meta_list = src_param_metadata.get(resolved_name) + if not src_meta_list and resolved_name.endswith("output_layer.weight"): + # Tied embeddings: the source shares the output projection with + # the input embedding, so it has no separate output_layer.weight. + # A pp>1 destination materializes one (embedding and output land + # on different stages), e.g. pp=1 (tied) -> pp=2. Source it from + # the embedding weight (same shape + vocab/TP shard); that tensor + # then feeds both the destination embedding and output_layer. + for emb_name in ("embedding.word_embeddings.weight", "word_embeddings.weight"): + src_meta_list = src_param_metadata.get(emb_name) + if src_meta_list: + break + if not src_meta_list: + raise RuntimeError( + f"Destination parameter '{resolved_name}' on rank {dst_rank} " + "not found in source model." + ) + # Choose a representative source metadata with DP round-robin balancing. + src_metadata = select_src_metadata_balanced(src_meta_list, dst_metadata, dst_rank) + sources = _determine_source_ranks_for_dst_param( + resolved_name, src_metadata, dst_metadata, dst_rank + ) + for src_rank, src_slice, dst_slice in sources: + task_id = next_task_id + next_task_id += 1 + yield task_id, dst_rank, src_rank, src_slice, dst_slice, src_metadata, dst_metadata + + +def _extract_module_metadata( + module, owner_rank, num_experts, rank_offset, rank_list_cache +) -> list[ParameterMetadata]: + """Metadata for a module's params and persistent buffers, or [] if None. + + Persistent buffers travel too so training state (e.g. MoE router expert_bias) + refits with the weights. + """ + if module is None: + return [] + pg = getattr(module, "pg_collection", None) + if pg is None: + raise ValueError("Module must have pg_collection") + layer_prefix_map = _build_layer_module_prefix_map(module) + return [ + extract_param_metadata( + p, + name, + owner_rank, + pg, + num_experts=num_experts, + layer_module_prefix_map=layer_prefix_map, + rank_offset=rank_offset, + _rank_list_cache=rank_list_cache, + ) + for name, p in named_refit_tensors(module) + ] + + +def index_metadata_rosters(gathered_pairs: list): + """Turn a rank-ordered list of ``(src_meta, dst_meta)`` (index == rank) into the + two rosters the plan builder consumes: dst params keyed by rank, and src params + keyed by resolved_name. The list may come from the all-gather, or be reassembled + in rank order as nodes are added, before calling build_plan_from_rosters. + """ + dst_param_metadata_by_rank: dict[int, dict[str, ParameterMetadata]] = {} + src_param_metadata: dict[str, list[ParameterMetadata]] = {} + for rank_id, (src_meta_list, dst_meta_list) in enumerate(gathered_pairs): + dst_param_metadata_by_rank[rank_id] = {m.resolved_name: m for m in dst_meta_list} + for metadata in src_meta_list: + src_param_metadata.setdefault(metadata.resolved_name, []).append(metadata) + return dst_param_metadata_by_rank, src_param_metadata + + +def build_plan_from_rosters( + dst_param_metadata_by_rank: dict[int, dict[str, ParameterMetadata]], + src_param_metadata: dict[str, list[ParameterMetadata]], + my_global_rank: int, +) -> ReshardPlan: + """Replay the deterministic global schedule and keep only this rank's ops. + + Pure and collective-free, so it can be tested or reused with preassembled + rosters without touching the process group. Live membership orchestration + is intentionally outside this module. + """ + my_plan = ReshardPlan([], []) + for ( + task_id, + dst_rank, + src_rank, + src_slice, + dst_slice, + src_metadata, + dst_metadata, + ) in _iter_global_transfer_ops(dst_param_metadata_by_rank, src_param_metadata): + if dst_rank == my_global_rank: + my_plan.recv_ops.append( + TransferOp( + param_name=dst_metadata.name, + peer_rank=src_rank, + is_send=False, + my_slice=dst_slice, + peer_slice=src_slice, + task_id=task_id, + ) + ) + if src_rank == my_global_rank: + my_plan.send_ops.append( + TransferOp( + param_name=src_metadata.name, + peer_rank=dst_rank, + is_send=True, + my_slice=src_slice, + peer_slice=dst_slice, + task_id=task_id, + ) + ) + + logger.info( + f"Rank {my_global_rank}: Built plan locally - {len(my_plan.recv_ops)} recvs, " + f"{len(my_plan.send_ops)} sends" + ) + return my_plan + + +def build_local_reshard_plan( src_module: torch.nn.Module, dst_module: torch.nn.Module, - num_experts: int = None, + num_experts: int | None = None, group=None, src_rank_offset: int = 0, dst_rank_offset: int = 0, ) -> ReshardPlan: """ - Centralized planning: Rank 0 builds complete plan for all ranks, then scatters. - - Supports None for src_module and/or dst_module to enable non-collocated mode: - - src_module=None: Rank doesn't have source model (destination-only) - - dst_module=None: Rank doesn't have destination model (source-only) - - Both provided: Rank has both models (collocated mode) - - Each rank provides metadata only for the models it owns, including parallel group - membership (tensor_parallel_group_ranks, expert_parallel_group_ranks, etc.). - This metadata is sufficient for rank 0 to build correct transfer plans without - requiring dummy models. + Build this rank's reshard plan locally: all-gather the parameter metadata, + replay the global schedule (see _iter_global_transfer_ops), and keep only the + ops where this rank is the sender or receiver. No rank-0 bottleneck and no + scatter, since sender and receiver derive matching task_ids from the same + metadata. + + The metadata gather (the one collective) and the plan build are split into + index_metadata_rosters + build_plan_from_rosters, so the deterministic build + can also be tested against preassembled rosters without a process group. + + src_module/dst_module may be None for non-collocated ranks (destination-only, + source-only, or idle). Each rank contributes metadata only for the models it + owns, including its parallel-group membership. """ - # Use group.rank() instead of dist.get_rank(group) to support cross-cluster - # ProcessGroups where members have independent default PGs (same default rank). + # group.rank()/size() (not dist.get_rank(group)) support cross-cluster PGs + # whose members have independent default PGs. my_global_rank = group.rank() if group is not None else dist.get_rank() world_size = group.size() if group is not None else dist.get_world_size() - # Shared cache for deduplicating rank lists across all metadata on this - # rank. Params sharing the same TP/DP/EP/PP groups will reference one - # list object, making pickle ~75% smaller for the gather. - _rank_list_cache: dict = {} - - def _extract_metadata(module, rank_offset): - """Extract per-parameter metadata from a module, or [] if module is None. - - Includes both ``nn.Parameter`` instances and persistent buffers — the - latter so that buffers carrying training state (e.g. MoE router - ``expert_bias``) travel with the weights during refit. - """ - if module is None: - return [] - pg = getattr(module, "pg_collection", None) - if pg is None: - raise ValueError("Module must have pg_collection") - layer_prefix_map = _build_layer_module_prefix_map(module) - return [ - extract_param_metadata( - p, - name, - my_global_rank, - pg, - num_experts=num_experts, - layer_module_prefix_map=layer_prefix_map, - rank_offset=rank_offset, - _rank_list_cache=_rank_list_cache, - ) - for name, p in named_refit_tensors(module) - ] - - my_src_metadata = _extract_metadata(src_module, src_rank_offset) - my_dst_metadata = _extract_metadata(dst_module, dst_rank_offset) - - # Gather (src, dst) tuples in one collective so we pay one pickle round-trip - # instead of two. Only rank 0 needs the full picture; other ranks just need - # their own plan from the later scatter. - gathered_pairs = [None] * world_size if my_global_rank == 0 else None - dist.gather_object((my_src_metadata, my_dst_metadata), gathered_pairs, group_dst=0, group=group) + # Dedup rank lists so params sharing a group reuse one list object; shrinks + # the pickled all-gather ~75%. + rank_list_cache: dict = {} + my_src_metadata = _extract_module_metadata( + src_module, my_global_rank, num_experts, src_rank_offset, rank_list_cache + ) + my_dst_metadata = _extract_module_metadata( + dst_module, my_global_rank, num_experts, dst_rank_offset, rank_list_cache + ) - # Free local metadata — no longer needed after gather. + # One all-gather gives every rank the full (src, dst) picture, replacing the + # gather-to-0 + scatter. + gathered_pairs = [None] * world_size + dist.all_gather_object(gathered_pairs, (my_src_metadata, my_dst_metadata), group=group) del my_src_metadata, my_dst_metadata - # Parameter to metadata maps keyed by resolved_name (only populated on rank 0) - dst_param_metadata_by_rank = {} - src_param_metadata: dict[str, list[ParameterMetadata]] = {} - - if my_global_rank == 0: - for rank_id, (src_meta_list, dst_meta_list) in enumerate(gathered_pairs): - dst_param_metadata_by_rank[rank_id] = {m.resolved_name: m for m in dst_meta_list} - for metadata in src_meta_list: - src_param_metadata.setdefault(metadata.resolved_name, []).append(metadata) - - # Free the raw gathered list — data is now in the indexed dicts. - del gathered_pairs - - # Build the plan on global rank 0 and broadcast to all ranks - if my_global_rank == 0: - plans_for_all_ranks = {r: ReshardPlan([], []) for r in range(world_size)} - # Global monotonically increasing ID for non-local transfers. - # This is shared between the corresponding send/recv ops so that - # NVSHMEM can build schedule. - next_task_id = 0 - - # Pipeline-parallel (PP) "mapping" is handled implicitly. - # Each rank contributes metadata only for the parameters it actually owns - # (i.e., the module partitioning for its PP stage). When PP sizes differ - # between source and destination, we don't compute an explicit stage-to-stage - # mapping here; instead, we iterate destination ranks and plan copies for the - # parameters present on those ranks. Any source rank that has the same logical - # parameter (matched by resolved_name) can serve as a sender (with DP balancing), - # and TP slicing is applied when applicable. - for dst_rank in range(world_size): - dst_rank_params = dst_param_metadata_by_rank.get(dst_rank, {}) - for resolved_name, dst_metadata in dst_rank_params.items(): - src_meta_list = src_param_metadata.get(resolved_name) - if not src_meta_list and resolved_name.endswith("output_layer.weight"): - # Tied embeddings: the source shares the output projection with - # the input embedding, so it has no separate output_layer.weight. - # A pp>1 destination materializes one (embedding and output land - # on different stages), e.g. pp=1 (tied) -> pp=2. Source it from - # the embedding weight (same shape + vocab/TP shard); that tensor - # then feeds both the destination embedding and output_layer. - for emb_name in ("embedding.word_embeddings.weight", "word_embeddings.weight"): - src_meta_list = src_param_metadata.get(emb_name) - if src_meta_list: - break - if not src_meta_list: - raise RuntimeError( - f"Destination parameter '{resolved_name}' on rank {dst_rank} " - "not found in source model." - ) - # Choose a representative source metadata with DP round-robin balancing - src_metadata = select_src_metadata_balanced(src_meta_list, dst_metadata, dst_rank) - sources = _determine_source_ranks_for_dst_param( - resolved_name, src_metadata, dst_metadata, dst_rank - ) - for src_rank, src_slice, dst_slice in sources: - task_id = next_task_id - next_task_id += 1 - - plans_for_all_ranks[dst_rank].recv_ops.append( - TransferOp( - param_name=dst_metadata.name, - peer_rank=src_rank, - is_send=False, - my_slice=dst_slice, - peer_slice=src_slice, - task_id=task_id, - ) - ) - plans_for_all_ranks[src_rank].send_ops.append( - TransferOp( - param_name=src_metadata.name, - peer_rank=dst_rank, - is_send=True, - my_slice=src_slice, - peer_slice=dst_slice, - task_id=task_id, - ) - ) - plans_list = [plans_for_all_ranks[r] for r in range(world_size)] - - # Free planning intermediates on rank 0 before the scatter. - del plans_for_all_ranks, dst_param_metadata_by_rank, src_param_metadata - else: - plans_list = None + dst_param_metadata_by_rank, src_param_metadata = index_metadata_rosters(gathered_pairs) + del gathered_pairs + return build_plan_from_rosters(dst_param_metadata_by_rank, src_param_metadata, my_global_rank) - # Scatter: each rank receives only its own plan (not all plans). - my_plan_list = [None] - torch.distributed.scatter_object_list(my_plan_list, plans_list, group_src=0, group=group) - my_plan = my_plan_list[0] - del plans_list # Free the full list on rank 0. - logger.info( - f"Rank {my_global_rank}: Received plan - {len(my_plan.recv_ops)} recvs, " - f"{len(my_plan.send_ops)} sends" +def build_centralized_reshard_plan( + src_module: torch.nn.Module, + dst_module: torch.nn.Module, + num_experts: int | None = None, + group=None, + src_rank_offset: int = 0, + dst_rank_offset: int = 0, +) -> ReshardPlan: + """Deprecated compatibility wrapper for :func:`build_local_reshard_plan`.""" + warnings.warn( + "build_centralized_reshard_plan is deprecated; use build_local_reshard_plan instead.", + DeprecationWarning, + stacklevel=2, + ) + return build_local_reshard_plan( + src_module, + dst_module, + num_experts=num_experts, + group=group, + src_rank_offset=src_rank_offset, + dst_rank_offset=dst_rank_offset, ) - - return my_plan diff --git a/megatron/core/resharding/refit.py b/megatron/core/resharding/refit.py index 36ba914d33b..a8e21826ae6 100644 --- a/megatron/core/resharding/refit.py +++ b/megatron/core/resharding/refit.py @@ -21,16 +21,17 @@ from megatron.core.models.common.language_module.language_module import LanguageModule from megatron.core.utils import unwrap_model -from . import build_centralized_reshard_plan, execute_reshard_plan +from . import build_local_reshard_plan, execute_reshard_plan from .copy_services.base import CopyService from .copy_services.gloo_copy_service import GlooCopyService from .copy_services.nccl_copy_service import NCCLCopyService +from .copy_services.nixl_copy_service import NixlCopyService from .copy_services.nvshmem_copy_service import NVSHMEMCopyService from .transforms import MXFP8ReshardTransform, ReshardTransform from .utils import invalidate_refit_tensor_cache, named_persistent_buffers # Supported refit backend names -RefitBackendName = Literal["nccl", "gloo", "nvshmem"] +RefitBackendName = Literal["nccl", "gloo", "nvshmem", "nixl"] @dataclass(frozen=True) @@ -44,6 +45,9 @@ class _PlanCacheKey: src_config: Optional[Tuple[int, int, int, int, int]] dst_config: Optional[Tuple[int, int, int, int, int]] num_experts: Optional[int] + # Adding inference nodes leaves the configs and offsets unchanged, so without + # world_size the stale pre-growth plan would be reused. + world_size: int = 0 # Rank offsets distinguish non-collocated configurations that would otherwise # share the same (rank, sizes, num_experts) tuple but route to different # global ranks. @@ -85,13 +89,16 @@ def _build_plan_cache_key( dst_rank_offset: int = 0, ) -> _PlanCacheKey: """Build cache key for reshard plan.""" - # group.rank() supports cross-cluster ProcessGroups. + # group.rank()/size() support cross-cluster ProcessGroups where members + # have independent default PGs. rank = group.rank() if group is not None else torch.distributed.get_rank() + world_size = group.size() if group is not None else torch.distributed.get_world_size() return _PlanCacheKey( rank=rank, src_config=_get_config_tuple(src_core), dst_config=_get_config_tuple(tgt_core), num_experts=num_experts, + world_size=world_size, src_rank_offset=src_rank_offset, dst_rank_offset=dst_rank_offset, ) @@ -109,7 +116,7 @@ def get_or_create_service(backend: RefitBackendName, group=None) -> CopyService: when swap_model_weights is called multiple times with the same backend. Args: - backend: Backend name ("nccl", "gloo", or "nvshmem"). + backend: Backend name ("nccl", "gloo", "nvshmem", or "nixl"). group: Optional process group for NCCL backend. """ if backend in _service_cache: @@ -121,6 +128,8 @@ def get_or_create_service(backend: RefitBackendName, group=None) -> CopyService: service = GlooCopyService(group=group) elif backend == "nvshmem": service = NVSHMEMCopyService(group=group) + elif backend == "nixl": + service = NixlCopyService(group=group) else: raise ValueError(f"Unknown backend '{backend}'") @@ -195,7 +204,8 @@ def _build_or_get_plan(src_core, tgt_core, num_experts, group, src_rank_offset, """Return the cached reshard plan, building it (collectively) if not yet cached. All participating ranks must call this simultaneously when the plan is not - yet cached, because build_centralized_reshard_plan uses collective communication. + yet cached, because build_local_reshard_plan uses collective communication + (an all_gather of parameter metadata). """ global _plan_cache cache_key = _build_plan_cache_key( @@ -207,7 +217,7 @@ def _build_or_get_plan(src_core, tgt_core, num_experts, group, src_rank_offset, dst_rank_offset=dst_rank_offset, ) if cache_key not in _plan_cache: - _plan_cache[cache_key] = build_centralized_reshard_plan( + _plan_cache[cache_key] = build_local_reshard_plan( src_core, tgt_core, num_experts=num_experts, diff --git a/megatron/core/ssm/gated_delta_net/__init__.py b/megatron/core/ssm/gated_delta_net/__init__.py new file mode 100644 index 00000000000..31147c6d31f --- /dev/null +++ b/megatron/core/ssm/gated_delta_net/__init__.py @@ -0,0 +1,37 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Gated Delta Net (GDN) family of layers. + +This package replaces the former ``megatron/core/ssm/gated_delta_net.py`` module +at the same import path; the names below preserve that module's public surface. +""" + +from megatron.core.ssm.gated_delta_net.common import ( + HAVE_FLA, + GatedDeltaNet, + GatedDeltaNetSubmodules, + _build_head_perm_for_split_sections, + _build_thd_cp_a2a_perm, + causal_conv1d, + chunk_gated_delta_rule, + get_parameter_local_cp_headwise, + l2norm, + tensor_a2a_cp2hp, + tensor_a2a_hp2cp, + torch_chunk_gated_delta_rule, +) + +__all__ = [ + "HAVE_FLA", + "GatedDeltaNet", + "GatedDeltaNetSubmodules", + "_build_head_perm_for_split_sections", + "_build_thd_cp_a2a_perm", + "causal_conv1d", + "chunk_gated_delta_rule", + "get_parameter_local_cp_headwise", + "l2norm", + "tensor_a2a_cp2hp", + "tensor_a2a_hp2cp", + "torch_chunk_gated_delta_rule", +] diff --git a/megatron/core/ssm/gated_delta_net.py b/megatron/core/ssm/gated_delta_net/common.py similarity index 100% rename from megatron/core/ssm/gated_delta_net.py rename to megatron/core/ssm/gated_delta_net/common.py diff --git a/megatron/core/ssm/gated_delta_net/gdn.py b/megatron/core/ssm/gated_delta_net/gdn.py new file mode 100644 index 00000000000..a94b17e0fae --- /dev/null +++ b/megatron/core/ssm/gated_delta_net/gdn.py @@ -0,0 +1,19 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2025, Songlin Yang, Jan Kautz, Ali Hatamizadeh. + +# Some of this code was adopted from https://github.com/huggingface/transformers +# This source code is licensed under the Apache license found in the +# LICENSE file in the root directory of this source tree. + +# pylint: disable=unused-import + +"""GatedDeltaNet layer. + +The full ``GatedDeltaNet`` implementation lives in +:mod:`megatron.core.ssm.gated_delta_net.common`; this module re-exports it so the +``megatron.core.ssm.gated_delta_net.gdn`` import path keeps working. +""" + +from megatron.core.ssm.gated_delta_net.common import GatedDeltaNet + +__all__ = ["GatedDeltaNet"] diff --git a/megatron/core/ssm/mamba_layer.py b/megatron/core/ssm/mamba_layer.py index d3b04e59c29..68d41c56a31 100644 --- a/megatron/core/ssm/mamba_layer.py +++ b/megatron/core/ssm/mamba_layer.py @@ -6,7 +6,7 @@ # LICENSE file in the root directory of this source tree. from dataclasses import dataclass, field -from typing import Dict, Optional, Protocol, Tuple, Union +from typing import Dict, Optional, Tuple, Union import torch from torch import Tensor @@ -21,18 +21,12 @@ from megatron.core.transformer.identity_op import IdentityOp from megatron.core.transformer.module import GraphableMegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module -from megatron.core.transformer.torch_norm import LayerNormInterface +from megatron.core.transformer.torch_norm import LayerNormBuilder from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.typed_torch import apply_module from megatron.core.utils import deprecate_inference_params -class LayerNormBuilder(Protocol): - """A protocol showing how MambaLayer expects to construct its LayerNorm.""" - - def __call__(self, config: TransformerConfig, hidden_size: int, /) -> LayerNormInterface: ... - - @dataclass class MambaLayerSubmodules: """ @@ -96,7 +90,11 @@ def __init__( pp_layer_offset=pp_layer_offset, name=(name + f".mixer") if name is not None else None, ) - self.norm = submodules.norm(self.config, self.config.hidden_size) + self.norm = submodules.norm( + config=self.config, + hidden_size=self.config.hidden_size, + eps=self.config.layernorm_epsilon, + ) self.mamba_bda = build_module(submodules.mamba_bda) self.bias_dropout_add_exec_handler = torch.enable_grad diff --git a/megatron/core/ssm/mamba_mixer.py b/megatron/core/ssm/mamba_mixer.py index c8b3ef583fe..24ac0d13964 100644 --- a/megatron/core/ssm/mamba_mixer.py +++ b/megatron/core/ssm/mamba_mixer.py @@ -8,7 +8,7 @@ import inspect import logging import math -from dataclasses import dataclass +from dataclasses import dataclass, replace from typing import List, Optional, Tuple, Union import torch @@ -26,9 +26,14 @@ from megatron.core.packed_seq_params import PackedSeqParams, resolve_cp_group from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.ssm.ops.causal_conv1d_triton import causal_conv1d_update +from megatron.core.ssm.ops.intermediate_extraction import ( + scatter_intermediate_conv, + scatter_intermediate_ssm, +) from megatron.core.ssm.ops.mamba_ssm import selective_state_update from megatron.core.ssm.utils import _split_tensor_factory from megatron.core.tensor_parallel import get_cuda_rng_tracker +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP from megatron.core.transformer import TransformerConfig from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec, build_module @@ -43,8 +48,14 @@ is_mamba_min_version, is_using_quantization_scales, log_single_rank, + make_tp_sharded_tensor_for_checkpoint, ) +if HAVE_GTP: + from megatron.core.tensor_parallel.gtp_api import is_gtp_param +else: + is_gtp_param = None + from .mamba_context_parallel import MambaContextParallel try: @@ -664,6 +675,7 @@ def _dynamic_inference_prefill( slot_allocator = context.mamba_slot_allocator intermediate_chunk_indices = metadata.intermediate_chunk_indices intermediate_abs_positions = metadata.intermediate_abs_positions + intermediate_real_count = metadata.intermediate_real_count intermediate_ssm_out = None intermediate_conv_out = None if slot_allocator is not None and mamba_layer_idx is not None: @@ -679,9 +691,9 @@ def _dynamic_inference_prefill( batch_indices=batch_indices, intermediate_chunk_indices=intermediate_chunk_indices, intermediate_abs_positions=intermediate_abs_positions, + intermediate_real_count=intermediate_real_count, intermediate_ssm_out=intermediate_ssm_out, intermediate_conv_out=intermediate_conv_out, - conv_gather_offsets=metadata.conv_gather_offsets, cu_chunk_seqlens=metadata.cu_chunk_seqlens, last_chunk_indices=metadata.last_chunk_indices, seq_idx_for_varlen=metadata.seq_idx_for_varlen, @@ -793,9 +805,9 @@ def _ssm_prefill( batch_indices: Optional[torch.Tensor] = None, intermediate_chunk_indices: Optional[torch.Tensor] = None, intermediate_abs_positions: Optional[torch.Tensor] = None, + intermediate_real_count: Optional[torch.Tensor] = None, intermediate_ssm_out: Optional[torch.Tensor] = None, intermediate_conv_out: Optional[torch.Tensor] = None, - conv_gather_offsets: Optional[torch.Tensor] = None, cu_chunk_seqlens: Optional[torch.Tensor] = None, last_chunk_indices: Optional[torch.Tensor] = None, seq_idx_for_varlen: Optional[torch.Tensor] = None, @@ -820,12 +832,13 @@ def _ssm_prefill( intermediate state extraction (fixed size, padded with 0). intermediate_abs_positions: Pre-allocated tensor of absolute token positions for conv state extraction (fixed size, padded with d_conv). + intermediate_real_count: int32[1] GPU tensor holding the number of + meaningful entries in the intermediate buffers this step. Read + inside the Triton scatter kernels so padded slots cost nothing. intermediate_ssm_out: Output buffer for extracted SSM states [max_intermediate_count, *ssm_shape]. intermediate_conv_out: Output buffer for extracted conv states [max_intermediate_count, *conv_shape]. - conv_gather_offsets: Constant tensor [-d_conv, ..., -1] for gathering - conv states. cu_chunk_seqlens: Precomputed chunk boundaries from MambaMetadata. last_chunk_indices: Precomputed last chunk index per sequence. seq_idx_for_varlen: Precomputed request ID per chunk. @@ -990,6 +1003,13 @@ def _ssm_prefill( chunk_starts = cu_chunk_seqlens[:-1] seq_idx_for_varlen = seq_idx[0, chunk_starts].contiguous() + # Extraction is enabled when the slot allocator wired buffers in via + # the caller. When enabled, the chunk scan returns its raw states so + # our Triton kernels do a fused gather+conditional-scatter directly, + # skipping the dense intermediate tensor and the padded-slot writes. + extract_intermediates = ( + intermediate_chunk_indices is not None and intermediate_ssm_out is not None + ) ssm_varlen_result = mamba_chunk_scan_combined_varlen( x=x, dt=dt, @@ -1009,43 +1029,44 @@ def _ssm_prefill( z=z if not self.rmsnorm else None, dt_bias=self.cp.get_dt_bias().float(), initial_states=initial_ssm_state, - return_intermediate_states=False, - intermediate_chunk_indices=intermediate_chunk_indices, + return_raw_states=extract_intermediates, dt_softplus=True, dt_limit=(0.0, float("inf")), state_dtype=ssm_state.dtype, ) - if intermediate_chunk_indices is not None: - ssm_varlen_states, intermediate_ssm_states = ssm_varlen_result + if extract_intermediates: + ssm_varlen_states, raw_ssm_states = ssm_varlen_result else: ssm_varlen_states = ssm_varlen_result - intermediate_ssm_states = None + raw_ssm_states = None y = y.unsqueeze(0) z = z.unsqueeze(0) tensor_masked_update(ssm_state, batch_indices, ssm_varlen_states) - # Write intermediate states to pre-allocated output buffers - # All tensor ops, no Python loops, fully CUDA graph compatible. - # The destination buffers are sized to the global max_intermediate_count - # but we only fill the per-graph-bucket prefix; readers consult - # per_request_intermediate_counts to know the real count. - if intermediate_chunk_indices is not None and intermediate_ssm_out is not None: - n = intermediate_ssm_states.shape[0] - intermediate_ssm_out[:n].copy_(intermediate_ssm_states) - - # Vectorized conv state extraction - # conv_gather_offsets: [d_conv] = [-d_conv, ..., -1] - gather_positions = ( - intermediate_abs_positions.unsqueeze(1).long() - + conv_gather_offsets.unsqueeze(0).long() - ) # [n, d_conv] - intermediate_conv = xBC_pre_conv[0, gather_positions, :] - # [n, d_conv, conv_dim] - intermediate_conv_out[:n].copy_(intermediate_conv.transpose(1, 2)) - # [n, conv_dim, d_conv] + if extract_intermediates: + # Fused gather+conditional-scatter for SSM: read row + # raw_ssm_states[chunk_indices[i]] into intermediate_ssm_out[i], + # only for i < real_count. + scatter_intermediate_ssm( + raw_ssm_states, + intermediate_chunk_indices, + intermediate_real_count, + intermediate_ssm_out, + ) + # Same pattern for conv: gather a length-d_conv window ending at + # abs_positions[i] (clamped into the valid token range) from + # xBC_pre_conv and scatter (transposed) into intermediate_conv_out[i], + # only for i < real_count. + scatter_intermediate_conv( + xBC_pre_conv, + intermediate_abs_positions, + intermediate_real_count, + intermediate_conv_out, + d_conv=intermediate_conv_out.shape[-1], + ) else: # Non-dynamic-batching path (static batching) initial_ssm_state = None @@ -1372,6 +1393,38 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): + 2 * self.ngroups_local_tp * self.d_state + self.nheads_local_tp ) + # Under GTP, in_proj.weight is GTP-sliced along axis 0. The [z|x|B|C|dt] split boundaries + # don't line up with GTP slice boundaries, so gather the shards back to TP-local size + # (strip the trailing pad rows from the gathered tail) and fall through to the same + # split path the non-GTP run uses — saved ckpt format matches a non-GTP run. + in_proj_gtp_remat_size = getattr(self.in_proj.weight, "gtp_remat_size", 1) + if in_proj_gtp_remat_size > 1 and HAVE_GTP and is_gtp_param(self.in_proj.weight): + gtp_remat_group = self.in_proj.weight.group + # in_proj.weight was already built at the sharded size by the submodule + # sharded_state_dict above — and, for native-FP8 GTP, dequantized to BF16 there + # (make_tp_sharded_tensor_for_checkpoint). Gather those (BF16) shards back to the + # full TP-local size so the [z|x|B|C|dt] split below matches a non-GTP run. + local = sharded_state_dict[f"{prefix}in_proj.weight"].data.contiguous() + gathered = torch.empty( + (local.shape[0] * in_proj_gtp_remat_size,) + local.shape[1:], + dtype=local.dtype, + device=local.device, + ) + torch.distributed.all_gather_into_tensor(gathered, local, group=gtp_remat_group) + if gathered.shape[0] != in_proj_dim: + gathered = gathered[:in_proj_dim].contiguous() + # Gathered weight is replicated across full dp_cp; replica_id needs only the DP slot. + dp_cp_rank = torch.distributed.get_rank(metadata['dp_cp_group']) + sharded_state_dict[f"{prefix}in_proj.weight"] = make_tp_sharded_tensor_for_checkpoint( + gathered, + f"{prefix}in_proj.weight", + tp_axis=0, + replica_id=(0, 0, dp_cp_rank), + prepend_offsets=sharded_offsets, + tp_group=self.tp_group, + dp_cp_group=metadata['dp_cp_group'], + ) + assert sharded_state_dict[f"{prefix}in_proj.weight"].data.size(0) == in_proj_dim, ( in_proj_dim, sharded_state_dict[f"{prefix}in_proj.weight"], @@ -1390,6 +1443,40 @@ def sharded_state_dict(self, prefix="", sharded_offsets=(), metadata=None): 0, ) + # GTP load-side inverse of the save-time all-gather (see + # docs/api-guide/core/generalized_tensor_parallel.md §3.3, in_proj + # note): the checkpoint stores the FULL TP-local in_proj.weight (pad stripped) under the + # 5 split keys [z|x|B|C|dt], so the default merge_fn cats them back to ``in_proj_dim`` + # rows with no padding. To reload into the live GTP param we must mirror init + # (``_gtp_slice_one_param``): F.pad the merged tensor with zeros up to + # ``gtp_local_size * gtp_remat_size``, then slice by ``gtp_rank``. GTP_remat_size=1 has no + # pad/slice. + if in_proj_gtp_remat_size > 1 and HAVE_GTP and is_gtp_param(self.in_proj.weight): + factory = sharded_state_dict[f"{prefix}in_proj.weight"] + gtp_local_rank = torch.distributed.get_rank(self.in_proj.weight.group) + gtp_local_size = self.in_proj.weight.data.size(0) + original_merge_fn = factory.merge_fn + + @torch.no_grad() + def _gtp_slice_after_cat( + sub_state_dict, + _orig=original_merge_fn, + _rank=gtp_local_rank, + _size=gtp_local_size, + _gtp_remat_size=in_proj_gtp_remat_size, + ): + full = _orig(sub_state_dict) + aligned_total = _size * _gtp_remat_size + pad_rows = aligned_total - full.shape[0] + if pad_rows > 0: + full = torch.nn.functional.pad(full, (0, 0, 0, pad_rows)) + start = _rank * _size + return full[start : start + _size].contiguous() + + sharded_state_dict[f"{prefix}in_proj.weight"] = replace( + factory, merge_fn=_gtp_slice_after_cat + ) + conv_dim = self.d_inner_local_tp + 2 * self.ngroups_local_tp * self.d_state assert sharded_state_dict[f"{prefix}conv1d_weight"].data.size(0) == conv_dim, ( conv_dim, diff --git a/megatron/core/ssm/ops/intermediate_extraction.py b/megatron/core/ssm/ops/intermediate_extraction.py new file mode 100644 index 00000000000..ec6552d798e --- /dev/null +++ b/megatron/core/ssm/ops/intermediate_extraction.py @@ -0,0 +1,227 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Fused gather + conditional-scatter kernels for Mamba intermediate-state +extraction used by prefix caching. + +These replace the two-step ``states[indices]`` (dense gather) + ``.copy_()`` +(scratch write) pattern with a single kernel that: + +1. Reads a runtime ``real_count`` from a fixed-address GPU tensor. +2. For each slot ``i < real_count``, gathers the source row indexed by the + per-slot index/position and writes it directly into the destination scratch. +3. For each slot ``i >= real_count``, returns immediately (no work, no write). + +This is CUDA-graph safe: the launch grid is sized at capture time to the maximum +possible slot count, but per-program execution is data-conditional on the +runtime ``real_count``, so padded slots cost almost nothing. +""" + +import torch +from torch import Tensor + +try: + import triton + import triton.language as tl + + HAVE_TRITON = True +except ImportError: + from unittest.mock import MagicMock + + from megatron.core.utils import null_decorator + + triton = MagicMock() + triton.jit = null_decorator + tl = MagicMock() + HAVE_TRITON = False + + +@triton.jit +def _scatter_intermediate_ssm_kernel( + states_ptr, # [num_chunks, state_flat]: source SSM states from the chunk scan + chunk_indices_ptr, # [max_count] int64: gather index per scratch slot + real_count_ptr, # int32[1]: number of meaningful slots this step + out_ptr, # [max_count, state_flat]: scratch destination (per layer) + state_flat, # tl.int32: product of (nheads, headdim, dstate) + states_row_stride, # tl.int64: stride between chunks in states_ptr + out_row_stride, # tl.int64: stride between slots in out_ptr + BLOCK_N: tl.constexpr, # columns per program +): + """Conditional gather+scatter for SSM intermediate states. + + Grid: ``(max_count, ceil(state_flat / BLOCK_N))``. Each program owns one + (slot, column-block) pair; programs with ``pid_slot >= real_count`` exit + immediately, so padded slots produce no HBM traffic. + """ + pid_slot = tl.program_id(0) + pid_col = tl.program_id(1) + + real_count = tl.load(real_count_ptr).to(tl.int32) + if pid_slot >= real_count: + return + + chunk_idx = tl.load(chunk_indices_ptr + pid_slot).to(tl.int64) + + col_offset = pid_col * BLOCK_N + cols = col_offset + tl.arange(0, BLOCK_N) + mask = cols < state_flat + + src = states_ptr + chunk_idx * states_row_stride + cols.to(tl.int64) + dst = out_ptr + pid_slot.to(tl.int64) * out_row_stride + cols.to(tl.int64) + + val = tl.load(src, mask=mask) + tl.store(dst, val, mask=mask) + + +def scatter_intermediate_ssm( + states: Tensor, chunk_indices: Tensor, real_count_gpu: Tensor, out: Tensor +) -> None: + """Gather rows of ``states`` at ``chunk_indices`` and scatter into ``out``, + for the first ``real_count_gpu`` slots only. + + Args: + states: ``(num_chunks, *ssm_state_shape)`` chunk-scan output for one layer. + chunk_indices: ``(max_count,)`` int64 per-slot gather index. + real_count_gpu: ``int32[1]`` GPU tensor with the runtime real count. + out: ``(max_count, *ssm_state_shape)`` destination scratch slice (one layer). + """ + assert states.is_cuda and chunk_indices.is_cuda and real_count_gpu.is_cuda and out.is_cuda + assert states.dim() >= 2, f"expected states to be at least 2D, got {states.shape}" + assert ( + out.shape[1:] == states.shape[1:] + ), f"per-slot shape mismatch: out {tuple(out.shape[1:])} vs states {tuple(states.shape[1:])}" + assert chunk_indices.dtype == torch.int64, chunk_indices.dtype + assert real_count_gpu.dtype == torch.int32 and real_count_gpu.numel() == 1 + + # The grid follows the per-step view length (chunk_indices.numel()), bounded + # above by out.shape[0] (the full scratch pool); programs past real_count + # exit immediately inside the kernel. + n_slots = int(chunk_indices.numel()) + if n_slots == 0: + return + if n_slots > out.shape[0]: + raise ValueError(f"chunk_indices length {n_slots} exceeds scratch capacity {out.shape[0]}") + + state_flat = 1 + for d in states.shape[1:]: + state_flat *= int(d) + + # Contiguity is required for the row-stride math; the engine pre-allocates + # these contiguous, so assert to fail loudly in tests rather than corrupt. + assert states.is_contiguous(), "states must be contiguous" + assert out.is_contiguous(), "out must be contiguous" + + BLOCK_N = 1024 + grid = (n_slots, triton.cdiv(state_flat, BLOCK_N)) + _scatter_intermediate_ssm_kernel[grid]( + states, + chunk_indices, + real_count_gpu, + out, + state_flat=state_flat, + states_row_stride=state_flat, + out_row_stride=state_flat, + BLOCK_N=BLOCK_N, + ) + + +@triton.jit +def _scatter_intermediate_conv_kernel( + src_ptr, # xBC_pre_conv batch-0 slice: addressed via strides + abs_positions_ptr, # [max_count] int32: extraction-window end position per slot + real_count_ptr, # int32[1]: meaningful slot count this step + out_ptr, # [max_count, conv_dim, d_conv]: scratch destination + seq_len, # tl.int32: bound for clamping + conv_dim, # tl.int32: feature dim + src_stride_s, # tl.int64: stride along seq_len + src_stride_c, # tl.int64: stride along conv_dim + out_slot_stride, # tl.int64: stride between slots in out_ptr (conv_dim * D_CONV) + D_CONV: tl.constexpr, + BLOCK_C: tl.constexpr, +): + """Conditional gather of a length-``D_CONV`` conv window per slot. + + Reads window ``[abs_pos - D_CONV, abs_pos)`` from ``src_ptr`` (clamped into + ``[0, seq_len)``) and writes it transposed into ``out[slot, :, :]`` of shape + ``(conv_dim, D_CONV)``. The transpose is folded into the write pattern to + match the slot allocator's storage layout. + """ + pid_slot = tl.program_id(0) + pid_c = tl.program_id(1) + + real_count = tl.load(real_count_ptr).to(tl.int32) + if pid_slot >= real_count: + return + + abs_pos = tl.load(abs_positions_ptr + pid_slot).to(tl.int32) + + c_offset = pid_c * BLOCK_C + c_idxs = c_offset + tl.arange(0, BLOCK_C) + c_mask = c_idxs < conv_dim + + # out[slot, c, j]: c outer, j inner (contiguous), so for fixed c the D_CONV + # values are stride-1. + slot_base = pid_slot.to(tl.int64) * out_slot_stride + + for j in tl.static_range(D_CONV): + p_raw = abs_pos - D_CONV + j + # Clamp into [0, seq_len - 1]. A no-op for real slots; defensive for any + # future caller that doesn't filter padded slots via real_count. + p = tl.maximum(0, tl.minimum(p_raw, seq_len - 1)) + + src = src_ptr + p.to(tl.int64) * src_stride_s + c_idxs.to(tl.int64) * src_stride_c + dst = out_ptr + slot_base + c_idxs.to(tl.int64) * D_CONV + j + + val = tl.load(src, mask=c_mask) + tl.store(dst, val, mask=c_mask) + + +def scatter_intermediate_conv( + src: Tensor, abs_positions: Tensor, real_count_gpu: Tensor, out: Tensor, d_conv: int +) -> None: + """Gather length-``d_conv`` conv windows from ``src`` at ``abs_positions`` and + scatter (transposed) into ``out``, for the first ``real_count_gpu`` slots only. + + Args: + src: ``(batch, seq_len, conv_dim)`` pre-conv xBC tensor. Batch is assumed + to be 1 (inference); only batch index 0 is read. + abs_positions: ``(max_count,)`` int32 extraction-window end position per + slot (window is ``[pos - d_conv, pos)``). + real_count_gpu: ``int32[1]`` GPU tensor with the runtime real count. + out: ``(max_count, conv_dim, d_conv)`` destination scratch slice (one layer). + d_conv: conv window length (constexpr in the kernel). + """ + assert src.is_cuda and abs_positions.is_cuda and real_count_gpu.is_cuda and out.is_cuda + assert src.dim() == 3, f"expected src (batch, seq_len, conv_dim), got {src.shape}" + assert src.shape[0] == 1, f"batch must be 1 for inference, got {src.shape[0]}" + assert abs_positions.dtype == torch.int32, abs_positions.dtype + assert real_count_gpu.dtype == torch.int32 and real_count_gpu.numel() == 1 + assert ( + out.dim() == 3 and out.shape[2] == d_conv + ), f"out shape {tuple(out.shape)} does not match (max_count, conv_dim, {d_conv})" + + _, conv_dim, _ = out.shape + n_slots = int(abs_positions.numel()) + if n_slots == 0: + return + if n_slots > out.shape[0]: + raise ValueError(f"abs_positions length {n_slots} exceeds scratch capacity {out.shape[0]}") + + seq_len = int(src.shape[1]) + src_stride_s = int(src.stride(1)) + src_stride_c = int(src.stride(2)) + + BLOCK_C = 128 + grid = (n_slots, triton.cdiv(conv_dim, BLOCK_C)) + _scatter_intermediate_conv_kernel[grid]( + src, + abs_positions, + real_count_gpu, + out, + seq_len=seq_len, + conv_dim=conv_dim, + src_stride_s=src_stride_s, + src_stride_c=src_stride_c, + out_slot_stride=conv_dim * d_conv, + D_CONV=d_conv, + BLOCK_C=BLOCK_C, + ) diff --git a/megatron/core/ssm/ops/ssd_combined.py b/megatron/core/ssm/ops/ssd_combined.py index 4fcee98b13e..ba43dbc295a 100644 --- a/megatron/core/ssm/ops/ssd_combined.py +++ b/megatron/core/ssm/ops/ssd_combined.py @@ -32,11 +32,10 @@ def _mamba_chunk_scan_combined_fwd( z=None, dt_bias=None, initial_states=None, - return_intermediate_states=False, seq_idx=None, cu_chunk_seqlens=None, last_chunk_indices=None, - intermediate_chunk_indices=None, + return_raw_states=False, dt_softplus=False, dt_limit=(0.0, float("inf")), state_dtype=None, @@ -149,15 +148,11 @@ def _mamba_chunk_scan_combined_fwd( initial_states=initial_states, ) - if return_intermediate_states: - return states - final_states = states[last_chunk_indices] - if intermediate_chunk_indices is not None: - intermediate_states = states[intermediate_chunk_indices] - return final_states, intermediate_states - else: - return final_states + if return_raw_states: + # Caller extracts any chunk-boundary states itself from the raw states. + return final_states, states + return final_states def mamba_chunk_scan_combined_varlen( @@ -177,8 +172,7 @@ def mamba_chunk_scan_combined_varlen( initial_states=None, dt_softplus=False, dt_limit=(0.0, float("inf")), - return_intermediate_states=False, - intermediate_chunk_indices=None, + return_raw_states=False, state_dtype=None, ): """ @@ -198,13 +192,13 @@ def mamba_chunk_scan_combined_varlen( dt_bias: (nheads,) initial_states: (batch, nheads, headdim, dstate) dt_softplus: Whether to apply softplus to dt - intermediate_chunk_indices: (N,) optional int64 tensor of chunk indices at which to - extract intermediate SSM states. When provided, returns (final_states, - intermediate_states) instead of just final_states. + return_raw_states: If True, returns ``(varlen_states, raw_states)`` where + ``raw_states`` is the full ``(nchunks, nheads, headdim, dstate)`` + chunk-boundary state tensor for the caller to extract from directly. state_dtype: The data type of the ssm state Return: varlen_states: (batch, nheads, headdim, dstate), or - (varlen_states, intermediate_states) if intermediate_chunk_indices is provided + (varlen_states, raw_states) if return_raw_states is True """ assert seq_idx is not None @@ -221,11 +215,10 @@ def mamba_chunk_scan_combined_varlen( z=z, dt_bias=dt_bias, initial_states=initial_states, - return_intermediate_states=return_intermediate_states, seq_idx=seq_idx, cu_chunk_seqlens=cu_chunk_seqlens, last_chunk_indices=last_chunk_indices, - intermediate_chunk_indices=intermediate_chunk_indices, + return_raw_states=return_raw_states, dt_softplus=dt_softplus, dt_limit=dt_limit, state_dtype=state_dtype, diff --git a/megatron/core/tensor_parallel/__init__.py b/megatron/core/tensor_parallel/__init__.py index 0852014a859..6147ae7b65d 100644 --- a/megatron/core/tensor_parallel/__init__.py +++ b/megatron/core/tensor_parallel/__init__.py @@ -10,8 +10,10 @@ ColumnParallelLinear, RowParallelLinear, VocabParallelEmbedding, + copy_gtp_attributes, copy_tensor_model_parallel_attributes, linear_with_grad_accumulation_and_async_allreduce, + param_is_not_gtp_duplicate, param_is_not_tensor_parallel_duplicate, set_defaults_if_not_set_tensor_model_parallel_attributes, set_tensor_model_parallel_attributes, @@ -58,7 +60,9 @@ "set_tensor_model_parallel_attributes", "set_defaults_if_not_set_tensor_model_parallel_attributes", "copy_tensor_model_parallel_attributes", + "copy_gtp_attributes", "param_is_not_tensor_parallel_duplicate", + "param_is_not_gtp_duplicate", "linear_with_grad_accumulation_and_async_allreduce", # mappings.py "copy_to_tensor_model_parallel_region", diff --git a/megatron/core/tensor_parallel/generalized_tensor_parallelism.py b/megatron/core/tensor_parallel/generalized_tensor_parallelism.py new file mode 100644 index 00000000000..08f17e54996 --- /dev/null +++ b/megatron/core/tensor_parallel/generalized_tensor_parallelism.py @@ -0,0 +1,2202 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Generalized Tensor Parallelism (GTP). + +GTP factors the weight-parallel domain into ``TP x GTP_remat`` (two orthogonal +sub-axes). A weight is sharded ``1/(TP * GTP_remat)`` along its partition dim: + +* The ``TP`` slice stays sharded through the GEMM — ordinary tensor parallelism; + the output is TP-sharded and reduced/gathered as usual. +* The ``GTP_remat`` slice is *rematerialized* just before the GEMM: only the + ``gtp_remat`` sub-group async all-gathers its part, so each rank's GEMM sees + the full TP slice. This trades extra all-gather traffic for ``1/GTP_remat`` + lower weight (and optimizer/grad) memory — ZeRO-3-on-the-weight on top of TP. + +``GTP_remat`` (the rematerialization sub-axis) has degree ``gtp_weight_remat_size``, +derived from ``--tensor-parallel-num-weight-shards``. + +Materialization uses a per-weight prefetch chain + ticket-based buffer cache +co-designed for CUDA graph capture/replay. Quantized AG (FP8 / MXFP8 / NVFP4) +composes with the sharding for compounding bandwidth reduction. + +See ``docs/api-guide/core/generalized_tensor_parallel.md`` for design and usage. +""" + +from __future__ import annotations + +import logging +import math +import re +from collections import defaultdict +from contextlib import contextmanager, nullcontext +from dataclasses import dataclass, field +from enum import Enum +from typing import Dict, List, Optional + +import torch +from packaging.version import Version + +from megatron.core.utils import log_single_rank + +logger = logging.getLogger(__name__) + +_GTP_TE_MIN_VERSION = Version("2.19.0.dev0") + +try: + import transformer_engine as te # noqa: F401 + + _te_version = Version(te.__version__) + if _te_version < _GTP_TE_MIN_VERSION: + raise ImportError( + f"megatron.core.tensor_parallel.gtp_api requires TransformerEngine " + f">= {_GTP_TE_MIN_VERSION} (found {_te_version})." + ) + + import transformer_engine_torch as tex + from transformer_engine.pytorch.constants import ( + MXFP8_BLOCK_SCALING_SIZE, + NVFP4_BLOCK_SCALING_SIZE, + ) + from transformer_engine.pytorch.distributed import ( + _NVFP4AllGatherAsyncHandle, + gather_along_first_dim, + in_fp8_activation_recompute_phase, + reduce_scatter_along_first_dim, + ) + from transformer_engine.pytorch.module.base import get_dummy_wgrad + from transformer_engine.pytorch.quantized_tensor import QuantizedTensor + from transformer_engine.pytorch.tensor import MXFP8TensorStorage, NVFP4TensorStorage + from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Quantizer + from transformer_engine.pytorch.utils import ( + nvtx_range_pop, + nvtx_range_push, + round_up_to_nearest_multiple, + ) + + HAVE_TE = True +except (ImportError, ModuleNotFoundError): + # TE unavailable/too old -> stub the TE-backed names so this module still imports, + # and flag GTP unusable via HAVE_TE (gtp_api.py surfaces this as HAVE_GTP=False). No + # GTP path runs without TE. The `annotations` future-import keeps the lone + # module-level TE reference (a dataclass field annotation) from being evaluated. + from unittest.mock import MagicMock + + te = tex = MagicMock() + MXFP8_BLOCK_SCALING_SIZE = NVFP4_BLOCK_SCALING_SIZE = None + _NVFP4AllGatherAsyncHandle = MagicMock() + gather_along_first_dim = reduce_scatter_along_first_dim = MagicMock() + in_fp8_activation_recompute_phase = MagicMock() + get_dummy_wgrad = MagicMock() + QuantizedTensor = MagicMock() + MXFP8TensorStorage = NVFP4TensorStorage = MagicMock() + MXFP8Quantizer = MagicMock() + nvtx_range_pop = nvtx_range_push = round_up_to_nearest_multiple = MagicMock() + HAVE_TE = False + + +class GTPChain(str, Enum): + """Prefetch chain identifier for n GTPShardedParam. + + GRAPHED — fwd/bwd captured by a CUDA graph (MLM _CudaGraphRunner). + UNGRAPHED — fwd/bwd runs eagerly. + + Chains never cross-link (prev_w/next_w stay within one chain). See + _classify_param_chain for the GRAPHED/UNGRAPHED rule. + """ + + GRAPHED = "GTP_graphed" + UNGRAPHED = "GTP_ungraphed" + + +# One-block-ahead prefetch for routed grouped experts (see docs §3.4 "Grouped-expert chains"): +# - own chain per weight role -> next_w links the SAME role of CONSECUTIVE MoE blocks, +# so each all-gather gets a whole block of runway instead of one GEMM; +# - fc1/fc2 stay SEPARATE (merging leaves fc2 a GEMM behind) but share ONE stream via +# _stream_key, so their gathers serialize instead of splitting bandwidth; +# - the "_graphed"/"_ungraphed" suffix keeps captured and eager ops off the same stream. +_GTP_REMAT_GROUPED_PREFIX = "GTP_remat_grouped_" +_GTP_REMAT_GROUPED_FC1 = f"{_GTP_REMAT_GROUPED_PREFIX}fc1" +_GTP_REMAT_GROUPED_FC2 = f"{_GTP_REMAT_GROUPED_PREFIX}fc2" + + +def _graphness_suffix(graphed: bool) -> str: + """Chain-id suffix encoding the CUDA-graph capture axis.""" + return "graphed" if graphed else "ungraphed" + + +def _chain_is_grouped(chain_id: str) -> bool: + """True for the per-role grouped-expert chains (``_GTP_REMAT_GROUPED_FC1`` / ``_FC2``).""" + return chain_id.startswith(_GTP_REMAT_GROUPED_PREFIX) + + +def _chain_is_graphed(chain_id: str) -> bool: + """True for any CUDA-graph-captured chain, including grouped fc1/fc2 chains. + + Every chain id ends in "graphed" or "ungraphed" (see ``_classify_param_chain``), so testing + the eager suffix is exact; a new chain id must keep that convention. + """ + return not chain_id.endswith("ungraphed") + + +# Active cuda_graph config, set by the integrator via set_cuda_graph_modules() before +# classify_gtp_chains(); consumed by _classify_param_chain. +_CUDA_GRAPH_MODULES: Optional[set] = None # scope tags, e.g. {"mamba","attn","moe_router"} +_MOE_SHARED_EXPERT_OVERLAP: bool = False # overlapped shared_experts can't be captured -> UNGRAPHED +_FULL_ITERATION: bool = False # whole step in one graph -> every param GRAPHED +# Empty cuda_graph_modules under per-layer CG = "graph every layer" == all tags present. +_ALL_LAYER_SCOPE_TAGS = frozenset({"mamba", "attn", "moe", "moe_router"}) + + +def set_cuda_graph_modules( + scope, moe_shared_expert_overlap: bool = False, cuda_graph_impl: str = "none" +): + """Record the active cuda_graph config for GTP chain classification. + + Called by MLM at init, before classify_gtp_chains(). ``cuda_graph_impl`` + disambiguates the empty-``scope`` cases: + - "none" -> CG disabled; all params UNGRAPHED. + - "full_iteration" -> whole step in one graph; all params GRAPHED. + - "local"/"transformer_engine" + empty scope -> graph every layer. + """ + global _CUDA_GRAPH_MODULES, _MOE_SHARED_EXPERT_OVERLAP, _FULL_ITERATION + _MOE_SHARED_EXPERT_OVERLAP = bool(moe_shared_expert_overlap) + _FULL_ITERATION = cuda_graph_impl == "full_iteration" + if _FULL_ITERATION: + _CUDA_GRAPH_MODULES = None # scope unused + elif cuda_graph_impl != "none" and not scope: + _CUDA_GRAPH_MODULES = set(_ALL_LAYER_SCOPE_TAGS) # graph every layer + else: + _CUDA_GRAPH_MODULES = set(scope) if scope else None + + +def _classify_param_chain(param_name: str) -> str: + """Map a GTPShardedParam name + active cuda_graph config to its chain id (a string). + + Full-iteration -> GRAPHED. Otherwise embedding/output_layer are UNGRAPHED, and + each layer kind (mixer, attention, shared experts) is GRAPHED iff its scope tag is in + cuda_graph_modules. Routed grouped experts (``.mlp.experts.``) are special-cased FIRST: + fc1/fc2 each go to their own homogeneous grouped chain (see ``_GTP_REMAT_GROUPED_FC1``), with + graphness following the same "moe" scope rule. + """ + n = param_name + G = GTPChain.GRAPHED.value + U = GTPChain.UNGRAPHED.value + + # Routed grouped experts: own homogeneous chain per weight-role (fc1/fc2) for one-block-ahead + # prefetch. Checked BEFORE the generic rules (".mlp.shared_experts." is a distinct substring, so + # shared experts never fall in here). + if ".mlp.experts." in n: + graphed = _FULL_ITERATION or bool(_CUDA_GRAPH_MODULES and "moe" in _CUDA_GRAPH_MODULES) + # The grouped split is an EAGER-only optimization: when MoE is captured, keep grouped + # weights in the plain GRAPHED chain so the cross-graph drain — wait_async_comms( + # GTPChain.GRAPHED.value) in cuda_graphs.py — still targets them by exact chain id. + if graphed: + return G + eager = _graphness_suffix(False) + if ".linear_fc1." in n: + return f"{_GTP_REMAT_GROUPED_FC1}_{eager}" + if ".linear_fc2." in n: + return f"{_GTP_REMAT_GROUPED_FC2}_{eager}" + # Unknown grouped role (e.g. single fused weight): keep it in the general chain. + return U + + if _FULL_ITERATION: + return G + + # embedding/output_layer live outside any per-layer CG runner. + if "embedding" in n or "output_layer" in n: + return U + + scope = _CUDA_GRAPH_MODULES + if not scope: # CG disabled + return U + + if ".mlp.shared_experts." in n: + if _MOE_SHARED_EXPERT_OVERLAP: + return U + return G if ("moe" in scope or "moe_router" in scope) else U + + if ".self_attention." in n or ".cross_attention." in n: + return G if "attn" in scope else U + + if ".mixer." in n: + return G if "mamba" in scope else U + + return U + + +def classify_gtp_chains(model) -> None: + """Walk model.named_parameters() and set chain_id on every GTPShardedParam. + + Call once at init, AFTER set_cuda_graph_modules() and BEFORE the first fwd of any + graphed param. Raises if an already-initialized param would be reclassified into a + different chain (its prev/next links are already wired into the wrong list). + """ + conflicts = [] + for name, param in model.named_parameters(): + if not is_gtp_param(param): + continue + target = _classify_param_chain(name) + if param.prefetch_initialized and param.chain_id != target: + conflicts.append((name, param.chain_id, target)) + continue + param.chain_id = target + + # Bwd-prefetch opt-out: embedding weight needs no bwd AG (wgrad is a + # scatter-add on sharded rows, input has no dgrad) — saves one collective. + if "embedding" in name: + param._need_weight_prefetch_bwd = False + if conflicts: + raise RuntimeError( + "classify_gtp_chains: the following params were already chain-initialized " + "with a different chain_id than the classifier would assign — this means " + "their chain links are already wired into the wrong list. Move classification " + "earlier in init. Conflicts: " + + ", ".join(f"{n}: {old!r}->{new!r}" for n, old, new in conflicts[:3]) + + ("..." if len(conflicts) > 3 else "") + ) + + +class GTPWeightState(Enum): + """State of a GTPShardedParam's AG / RS lifecycle (debug / stale-read guard).""" + + NONE = "NONE" # Sharded, no pending operation + ASYNC_WAIT = "ASYNC_WAIT" # Async all-gather in progress + DATA_READY = "DATA_READY" # Async all-gather complete, result in cache + DATA_READY_SYNC = "DATA_READY_SYNC" # Sync all-gather complete, result in cache + + +# Global GTP buffer cache (persists across clear(); never set to None after creation). +_GTP_CACHE = None +_GTP_PARAMS = [] + +# Global set of GTPShardedParam with in-flight async comms (AG or RS). +_inflight_comm_params: set = set() +_AG_STREAMS: Dict[str, torch.cuda.Stream] = {} +_RS_STREAMS: Dict[str, torch.cuda.Stream] = {} + +# Wgrad input buffer pool, keyed by (shape, dtype). UNGRAPHED-only: GRAPHED +# wgrad bufs need address stability for CG replay and are not pool-recycled. +_wgrad_buf_pool: Dict[tuple, list] = {} + +# Double-buffering for the grouped one-block-ahead chains (docs §3.4): +# - the weight cache shares ONE buffer per (shape, dtype, expert_idx) -> safe only while at +# most one same-key weight is live; +# - one-block-ahead keeps blocks N and N+1 live at once, so without a tiebreak the prefetch +# would clobber the weight the running GEMM is still reading; +# - fold a chain-position parity (0,1,0,1...) into the cache key -> consecutive blocks +# alternate between exactly TWO buffers. +# Parity is assigned on first cache-key use, which happens in forward (= chain) order, so this +# per-(shape, expert_idx, chain_id) counter yields the alternating sequence. Cleared by +# reset_gtp_state so a rebuilt model restarts numbering. +_GTP_GROUPED_BUF_PARITY_COUNTER: Dict[tuple, int] = {} + + +def _wgrad_pool_get(shape: tuple, dtype: torch.dtype, device) -> torch.Tensor: + """Get a pool buffer or allocate fresh, tagged so _wgrad_pool_put accepts only + pool-owned buffers (other callers fall through to the caching allocator on release).""" + key = (shape, dtype) + pool = _wgrad_buf_pool.get(key) + if pool: + buf = pool.pop() + else: + buf = torch.empty(shape, dtype=dtype, device=device, requires_grad=False) + buf._from_gtp_wgrad_pool = True + return buf + + +def _wgrad_pool_put(buf: torch.Tensor): + """Return a pool-owned buffer for reuse (no-op for untagged buffers; see + _wgrad_pool_get).""" + if not getattr(buf, "_from_gtp_wgrad_pool", False): + return + key = (tuple(buf.shape), buf.dtype) + if key not in _wgrad_buf_pool: + _wgrad_buf_pool[key] = [] + _wgrad_buf_pool[key].append(buf) + + +def _stream_key(chain_id: str, group) -> tuple: + """Key for the per-(chain, group) AG/RS stream dicts. + + Partitioned on two axes: chain_id (captured GRAPHED vs eager UNGRAPHED ops must not + share a stream) and group (independent NCCL, e.g. GTP_remat vs EGTP_remat, no serialization). + + Grouped fc1/fc2 are separate chains but must share ONE stream, so their gathers serialize + instead of splitting bandwidth: drop the role from the key, keep the capture suffix. + """ + if _chain_is_grouped(chain_id): + chain_id = _GTP_REMAT_GROUPED_PREFIX + _graphness_suffix(_chain_is_graphed(chain_id)) + return (chain_id, id(group) if group is not None else 0) + + +def get_ag_stream(chain_id: str = GTPChain.GRAPHED.value, group=None) -> torch.cuda.Stream: + """Return the GTP all-gather stream for (chain_id, group). See _stream_key.""" + key = _stream_key(chain_id, group) + if key not in _AG_STREAMS: + _AG_STREAMS[key] = torch.cuda.Stream() + return _AG_STREAMS[key] + + +def get_rs_stream(chain_id: str = GTPChain.GRAPHED.value, group=None) -> torch.cuda.Stream: + """Return the GTP reduce-scatter stream for (chain_id, group). See _stream_key.""" + key = _stream_key(chain_id, group) + if key not in _RS_STREAMS: + _RS_STREAMS[key] = torch.cuda.Stream() + return _RS_STREAMS[key] + + +def wait_for_gtp_grad_reduction_on_current_stream() -> None: + """Fence the current stream against all GTP backward grad work before the DP gradient sync. + + Drains the eager AG/RS side streams, then waits on each CG runner's replay stream + (its tail = captured Phase 2 main_grad.add_). No-op when GTP is inactive. + """ + wait_async_comms() + cur = torch.cuda.current_stream() + for s in _AG_STREAMS.values(): + cur.wait_stream(s) + for s in _RS_STREAMS.values(): + cur.wait_stream(s) + # Local import: cuda_graphs imports this module, so a module-level import would be circular. + from megatron.core.transformer.cuda_graphs import get_gtp_runner_streams + + for s in get_gtp_runner_streams(): + cur.wait_stream(s) + + +@dataclass +class GTPRematConfig: + """Global configuration for Generalized Tensor Parallelism (weight remat).""" + + pad_for_alignment: int = 16 + check_param_states: bool = False + weight_prefetch: bool = True + # True (default): non-chain-head wgrad RS is async_op=True and finalizes + # (handle.wait + main_grad.add_) in a later bwd's cascade walk, overlapping RS with + # compute. False: every wgrad RS is synchronous + inline (no overlap). + async_reduction: bool = True + # Mirrors config.calculate_per_token_loss. When True, DDP applies NO 1/dp pre-scaling + # (gradient_scaling_factor=1.0) and finalize_model_grads normalizes every gradient by + # 1/total_global_tokens instead. In that mode the gtp_remat axis must be SUM-reduced (plain + # reduce-scatter, like DP), NOT mean-reduced — a 1/gtp mean would double-count the + # normalization. When False, the gtp_remat reduce-scatter takes the MEAN so it composes with + # DDP's 1/replicate scaling to yield the full (replicate x gtp) mean. + calculate_per_token_loss: bool = False + + +GTP_CONFIG = GTPRematConfig() + + +def update_gtp_config(**kwargs): + """Update the global GTP configuration.""" + for key, value in kwargs.items(): + if not hasattr(GTP_CONFIG, key): + raise ValueError(f"Unknown GTP config option: {key}") + setattr(GTP_CONFIG, key, value) + + +def tag_gtp_params_with_names(model): + """Populate _debug_name on every GTPShardedParam with its full dotted parameter name. + + Call once after model construction so the linking log prints human-readable names + instead of raw tensor ids. + """ + for name, param in model.named_parameters(): + if is_gtp_param(param): + param._debug_name = name + + +def configure_gtp_remat_from_recipe( + *, fp4=False, fp8_recipe=None, fp8=False, calculate_per_token_loss=False +): + """ + Configure GTP weight-remat (padding + loss reduction) from the quantization recipe. + Must be called once BEFORE model construction. + """ + # gtp_remat grad reduction SUMs (not means) the gtp_remat axis under per-token-loss. + # check_param_states=False: GTP buffer reuse (notably under CUDA-graph capture) trips the + # param-state debug asserts, so keep them off for GTP runs. + update_gtp_config(calculate_per_token_loss=calculate_per_token_loss, check_param_states=False) + if fp4: + update_gtp_config(pad_for_alignment=16) + elif fp8_recipe == "mxfp8": + update_gtp_config(pad_for_alignment=32) + elif fp8: + update_gtp_config(pad_for_alignment=16) + + if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0: + logger.info("> GTP_remat enabled. %s", GTP_CONFIG) + + +def classify_gtp_remat_chains( + model, *, cuda_graph_modules=None, moe_shared_expert_overlap=False, cuda_graph_impl="none" +): + """ + Tag and classify every GTP param's prefetch chain (GRAPHED vs UNGRAPHED). + Must be called once AFTER model build + DDP wrap and before the first forward (which + lazily builds chain links). + """ + cg_modules = ( + {getattr(s, "name", str(s)) for s in cuda_graph_modules} if cuda_graph_modules else None + ) + set_cuda_graph_modules( + cg_modules, + moe_shared_expert_overlap=moe_shared_expert_overlap, + cuda_graph_impl=cuda_graph_impl, + ) + # Clear stale process-global chain state so a rebuilt model starts fresh. + reset_gtp_state() + for model_module in model if isinstance(model, list) else [model]: + tag_gtp_params_with_names(model_module) + classify_gtp_chains(model_module) + + +def gtp_remat_shard_dim0(dim0, gtp_remat_group): + """Return ``(shard_dim0, pad_length)`` for allocating a dim-0 GTP weight-remat shard.""" + gtp_remat_size = gtp_remat_group.size() + if GTP_CONFIG.pad_for_alignment > 0: + alignment = GTP_CONFIG.pad_for_alignment * gtp_remat_size + pad_length = (alignment - dim0 % alignment) % alignment + else: + assert dim0 % gtp_remat_size == 0, ( + f"gtp_remat_shard_dim0: dim0={dim0} not divisible by gtp_remat_size={gtp_remat_size}. " + "Enable padding (GTP_CONFIG.pad_for_alignment > 0) or make dim-0 a multiple of the " + "GTP group size." + ) + pad_length = 0 + padded = dim0 + pad_length + return padded // gtp_remat_size, pad_length + + +def _gtp_slice_one_param(param, gtp_remat_group, *, name=""): + """Pad + slice a full-size BF16 weight to this rank's GTP shard. + + Caller attaches GTP attrs (see _gtp_attach_attrs). On the legacy post-init path under + fp8_model_init, tensor may be a QuantizedTensor — F.pad dequantizes it before slicing. + """ + gtp_remat_size = gtp_remat_group.size() + gtp_rank = gtp_remat_group.rank() + tensor = param.data + + if GTP_CONFIG.pad_for_alignment > 0: + # Pad before slicing so shards stay alignment-divisible and padding + # ends up contiguous at the tail of the gathered result. + alignment = GTP_CONFIG.pad_for_alignment * gtp_remat_size + dim0 = tensor.shape[0] + pad_length = (alignment - dim0 % alignment) % alignment + if pad_length > 0: + tensor = torch.nn.functional.pad(tensor, (0, 0, 0, pad_length)) + else: + # No-pad mode: dim-0 must divide gtp_remat_size or AG output loses tail rows. + assert tensor.shape[0] % gtp_remat_size == 0, ( + f"_gtp_slice_one_param: {name}.shape[0]={tensor.shape[0]} is not " + f"divisible by gtp_remat_size={gtp_remat_size}. Either enable padding by " + "setting GTP_CONFIG.pad_for_alignment > 0, or ensure the weight's " + "dim-0 is a multiple of the GTP group size." + ) + pad_length = 0 + + shard_size = tensor.shape[0] // gtp_remat_size + shard = tensor[gtp_rank * shard_size : (gtp_rank + 1) * shard_size] + gtp_shard = GTPShardedParam(shard.clone()) + gtp_shard.pad_length = pad_length + # Preserve the source weight's TP attributes (dropped when wrapping into GTPShardedParam), + # so param_is_not_tensor_parallel_duplicate still classifies it without GTP-specific code. + from megatron.core.tensor_parallel import copy_tensor_model_parallel_attributes + + copy_tensor_model_parallel_attributes(gtp_shard, param) + return gtp_shard + + +def _gtp_attach_attrs(gtp_shard, gtp_remat_group, *, is_grouped=False, expert_idx=0): + """Attach group / gtp_remat_size / routed-expert tags and register in _GTP_PARAMS. + + Separate from _gtp_slice_one_param so attrs land on the post-quantize param (when + quantize fires between slice and attach). + """ + # DistributedWeight requires implementers stay torch.Tensor subclasses; enforce at construction. + assert isinstance(gtp_shard, torch.Tensor), ( + "GTP param must remain a torch.Tensor subclass (DistributedWeight requirement); got " + f"{type(gtp_shard).__name__}." + ) + if is_grouped: + gtp_shard.expert_idx = expert_idx + gtp_shard.is_routed_expert = True + # Default to UNGRAPHED; classify_gtp_chains() reclassifies based on the + # cuda_graph_modules at init time. + gtp_shard.chain_id = GTPChain.UNGRAPHED.value + gtp_shard.group = gtp_remat_group + gtp_shard.gtp_remat_size = gtp_remat_group.size() + global _GTP_PARAMS + _GTP_PARAMS.append(gtp_shard) + + +def _gtp_wrap_bf16_shard(module, name, param): + """Re-register a BF16 pre-sharded weight as a :class:`GTPShardedParam`. + + The weight already IS this rank's shard (built pre-sharded), so — unlike the post-init path + :func:`_gtp_slice_one_param`, which slices a full weight — this only wraps it, no slicing. + Returns the new param (also swapped into the module). + """ + from megatron.core.tensor_parallel import copy_tensor_model_parallel_attributes + + gtp_shard = GTPShardedParam(param.data) + copy_tensor_model_parallel_attributes(gtp_shard, param) + delattr(module, name) + module._parameters[name] = gtp_shard + return gtp_shard + + +def _gtp_reclass_native_fp8_shard(param): + """Reclass a native-FP8 pre-sharded weight into a GTP subclass in place (buffer-resident). + + The dynamic ``GTP_`` subclass carries GTPShardedParam's gather/RS methods while + keeping ``is_float8tensor`` True (so DDP/distopt keep it buffer-resident) and TE's forward can + call ``weight.all_gather_and_prefetch()``. The param IS this rank's FP8 shard (``quantized`` is + itself; never re-quantized). Returns the same (mutated) param. + """ + # Preserve the param's own _quantizer (used by TE ops + the optimizer copy_->quantize_ update); + # _init_gtp_runtime_attrs clears it, so stash/restore. + native_quantizer = getattr(param, "_quantizer", None) + param.__class__ = _gtp_native_fp8_subclass(type(param)) + _init_gtp_runtime_attrs(param) + param._gtp_native_fp8 = True + param._quantizer = native_quantizer + # Gather uses a SEPARATE quantizer copy for its per-direction set_usage; reusing the param's own + # would leave rowwise=False after a bwd gather, freezing the rowwise data the forward reads. The + # copy also keeps MXFP8 scales compact for byte-concat scale all-gather. + gather_q = None + if native_quantizer is not None: + gather_q = native_quantizer.copy() + gather_q.internal = False + gather_q.optimize_for_gemm = not isinstance(gather_q, MXFP8Quantizer) + param._gtp_gather_quantizer = gather_q + param.quantized = param + return param + + +def attach_gtp_to_presharded_module(module, gtp_remat_group, pad_length, is_grouped=False): + """Turn each pre-sharded weight into a GTP param (FP8/BF16) and attach GTP wiring.""" + # GTP shards per-expert weight0..weight{num_gemms-1}; a coalesced single weight has no sibling + # shards to attach, so reject it here (once, at setup) instead of silently attaching nothing. + if is_grouped: + assert not getattr( + module, "single_grouped_weight", False + ), f"GTP grouped module {type(module).__name__} requires single_grouped_weight=False." + # Use the module's weight_names if it declares them; otherwise create them (grouped modules + # expose per-expert weight0..weight{num_gemms-1}, non-grouped a single "weight"). + weight_names = getattr(module, "weight_names", None) + if not weight_names: + weight_names = ( + [f"weight{idx}" for idx in range(module.num_gemms)] if is_grouped else ["weight"] + ) + new_weights = [] + for idx, name in enumerate(weight_names): + param = getattr(module, name, None) + if param is None or is_gtp_param(param): + continue + if isinstance(param, QuantizedTensor): + gtp_param = _gtp_reclass_native_fp8_shard(param) + else: + gtp_param = _gtp_wrap_bf16_shard(module, name, param) + gtp_param.pad_length = pad_length + _gtp_attach_attrs(gtp_param, gtp_remat_group, is_grouped=is_grouped, expert_idx=idx) + new_weights.append(gtp_param) + if is_grouped and new_weights: + new_weights[0].weight_list = new_weights + + +# Cache of dynamic ``GTP_`` subclasses, keyed by the FP8 base class. +_GTP_NATIVE_FP8_SUBCLASSES: Dict[type, type] = {} + +# GTPShardedParam members NOT copied into the dynamic ``GTP_`` subclass (and why): +# - __new__ / __init__: construction hooks — we only *reclass* an existing FP8 instance, never +# construct one; GTP attrs are set afterwards by _init_gtp_runtime_attrs. +# - __torch_function__: keep the FP8 tensor's OWN tensor dispatch, not GTPShardedParam's. +# - __dict__ / __weakref__ / __module__ / __doc__ / __qualname__ / __slots__: per-class machinery +# for GTPShardedParam itself; copying it would corrupt the new subclass's identity/layout. +_GTP_SUBCLASS_SKIP = frozenset( + { + "__new__", + "__init__", + "__torch_function__", + "__dict__", + "__weakref__", + "__module__", + "__doc__", + "__qualname__", + "__slots__", + } +) + + +def _gtp_native_fp8_subclass(base_cls: type) -> type: + """Cached ``base_cls`` subclass + GTPShardedParam's methods (isinstance(base_cls) kept True). + + CRITICAL: skip names already in the FP8 MRO — a GTPShardedParam method shadowing one TE needs + (e.g. get_data_tensors) silently freezes optimizer updates. + """ + sub = _GTP_NATIVE_FP8_SUBCLASSES.get(base_cls) + if sub is None: + base_mro_names = set() + for klass in base_cls.__mro__: + base_mro_names.update(vars(klass).keys()) + ns = { + k: v + for k, v in vars(GTPShardedParam).items() + if k not in _GTP_SUBCLASS_SKIP and k not in base_mro_names + } + # Share the (class-level) prefetch chain-state dicts with GTPShardedParam. + ns["_chain_state"] = GTPShardedParam._chain_state + ns["_recompute_chain_state"] = GTPShardedParam._recompute_chain_state + sub = type(f"GTP_{base_cls.__name__}", (base_cls,), ns) + _GTP_NATIVE_FP8_SUBCLASSES[base_cls] = sub + return sub + + +def is_gtp_param(param) -> bool: + """True if ``param`` is a GTP weight-remat shard (BF16 or native-FP8).""" + return getattr(param, "is_gtp_weight_remat", False) + + +def dequantize_gtp_native_fp8(param): + """Dequantize a native-FP8 GTP param to a plain BF16 tensor (used at the checkpoint boundary). + + TE's ``tex.dequantize`` dispatches on the *exact* FP8 class and rejects our dynamic + ``GTP_`` subclass, so restore the base FP8 class for the call and reclass after. + """ + from megatron.core.fp8_utils import dequantize_fp8_tensor + + sub_cls = type(param) + base_cls = sub_cls.__mro__[1] # _gtp_native_fp8_subclass builds type("GTP_X", (base_cls,), ...) + param.__class__ = base_cls + try: + return dequantize_fp8_tensor(param) + finally: + param.__class__ = sub_cls + + +@contextmanager +def gtp_native_fp8_load_context(module): + """Restore the base FP8 class on native-FP8 GTP params under ``module`` for a load copy. + + ``load_state_dict`` does ``param.copy_(bf16)`` -> TE ``convert_and_update_tensor``, whose + ``IsMXFP8Tensor`` C++ check rejects our dynamic subclass (the load-side twin of + :func:`dequantize_gtp_native_fp8`). Presenting the base class lets TE re-quantize into the FP8 + storage; instance attrs live in ``__dict__`` and survive the swap, so the GTP surface persists. + """ + from megatron.core.fp8_utils import is_float8tensor + + swapped = [] + for param in module.parameters(recurse=True): + if is_gtp_param(param) and is_float8tensor(param): + sub_cls = type(param) + base_cls = sub_cls.__mro__[1] + if base_cls is not sub_cls: + param.__class__ = base_cls + swapped.append((param, sub_cls)) + try: + yield + finally: + for param, sub_cls in swapped: + param.__class__ = sub_cls + + +def wrap_module_params_gtp(module, weight_names, gtp_remat_group, is_grouped=None): + """Shard and re-register module params as GTPShardedParam (post-init slice). + + Called post-init for Megatron-style local modules (ColumnParallelLinear, etc.), which build + the full weight and slice it here. TE modules do NOT use this path — they are constructed + already-shard-sized (GTP-agnostic init) and wired via :func:`attach_gtp_to_presharded_module`. + Params that are already GTP are skipped. + """ + if gtp_remat_group.size() == 1: + return + + for idx, name in enumerate(weight_names): + param = getattr(module, name, None) + if param is None: + continue + + # Already a GTP param (TE-side slice, or native-FP8 attach) — skip. + if is_gtp_param(param): + continue + + # delete the original parameter, which will be replaced by an GTP sharded one + delattr(module, name) + gtp_shard = _gtp_slice_one_param(param, gtp_remat_group, name=name) + del param + _gtp_attach_attrs(gtp_shard, gtp_remat_group, is_grouped=bool(is_grouped), expert_idx=idx) + # register the newly sharded param back to the module + module._parameters[name] = gtp_shard + + if is_grouped: + allweights = [getattr(module, name) for name in weight_names] + allweights[0].weight_list = allweights + + +class GTPShardHandle: + """Wrapper around a ``dist`` async-work handle for a GTP AG / RS. + + Tracks the participating shards so the wait-site can transition their GTPWeightState + and prune the param from _inflight_comm_params when the collective completes. + """ + + def __init__(self, handle, gtp_shards, reduce_scatter=False): + self.handle = handle + self.gtp_shards = gtp_shards + self.reduce_scatter = reduce_scatter + _inflight_comm_params.add(gtp_shards[0]) + + def wait(self): + """Wait on the underlying NCCL work and update the shards' state.""" + if self.handle is not None: + self.handle.wait() + self.handle = None # Release NCCL Work and its C++ tensor references promptly + if GTP_CONFIG.check_param_states: + for w in self.gtp_shards: + if self.reduce_scatter: + w._set_rs_state(GTPWeightState.DATA_READY) + else: + w._set_state(GTPWeightState.DATA_READY) + + _inflight_comm_params.discard(self.gtp_shards[0]) + + +def _init_gtp_runtime_attrs(obj): + """Initialize the full GTP runtime-state attribute surface on ``obj``. + + Shared by :meth:`GTPShardedParam.__init__` (legacy BF16 slice path) and + :func:`attach_gtp_to_presharded_module` (native-FP8 reclass path), so both param-class + representations carry identical state. chain_id/group are set by the caller afterward. + """ + # Canonical flag — also set on distopt's main_param copy so both kinds + # of param can be classified via a single attribute check. + obj.is_gtp_weight_remat = True + # all gather + obj.state = GTPWeightState.NONE + obj._ag_ticket_fwd = None + obj._ag_ticket_bwd = None + obj._prefetch_handle = None + obj._need_weight_prefetch = True + # Per-direction prefetch opt-outs (default True). The embedding weight needs no bwd AG + # (wgrad is a token-indexed scatter-add, input non-differentiable). classify_gtp_chains() + # sets this False for embedding.word_embeddings.weight. + obj._need_weight_prefetch_bwd = True + obj.ag_event = torch.cuda.Event(external=True) + # DDP backward hook (set by register_grad_accum_hook); invoked after + # the wgrad RS accumulation completes (Graphed.backward / chain cascade). + obj._grad_accum_hook = None + # Quantization. For native-FP8 GTP the reclass path overwrites _quantizer with the tensor's + # own MXFP8 quantizer and points quantized at self; BF16 GTP leaves both unset. + obj._quantizer = None + obj.quantized = None + # Prefetching linked list + obj.prefetch_initialized = False + obj.next_w = None + obj.prev_w = None + # Recompute-forward prefetch chain: a SEPARATE chain (own slot) for weights re-gathered + # rowwise during an activation-recompute forward in backward. Distinct from the + # state/_prefetch_handle/ag_event above so it never clobbers the concurrent columnwise + # dgrad lifecycle. Self-populates from the first backward's recompute gathers. + obj._recompute_initialized = False + obj._recompute_next = None + obj._recompute_prev = None + obj._recompute_prefetch_handle = None + obj._recompute_ag_event = torch.cuda.Event(external=True) + obj._recompute_already_drained = False + # Chain identity (GRAPHED/UNGRAPHED). Defaults to UNGRAPHED; classify_gtp_chains(model) + # walks the model at init (after set_cuda_graph_modules) and reclassifies on param name + + # active cuda_graph_modules. + obj.chain_id = GTPChain.UNGRAPHED.value + # Grouped gemm + obj.is_routed_expert = False + obj.expert_idx = None + obj.group = None + obj.weight_list = None + # Reduce-scatter state (set during wgrad_reduce_scatter) + obj.rs_state = GTPWeightState.NONE + obj._wgrad_rs_handle = None + obj.rs_event = torch.cuda.Event(external=True) + obj._rs_ticket = None + # Padding + obj.pad_length = 0 + # Debug + obj._debug_name = "" + # Hot-path caches (populated lazily on first use). chain_id/group are + # set after init, so we can't resolve streams eagerly here. + obj._cached_ag_stream = None + obj._cached_rs_stream = None + obj._cached_dtypes = None + obj._cached_gtp_remat_group = None + + +class GTPShardedParam(torch.nn.Parameter): + """A weight parameter sharded 1/N across a GTP process group. + + Materialized on-demand via async all-gather and gradient-reduced via reduce-scatter. + Carries its own prefetch-chain wiring (prev_w/next_w), per-chain state, AG/RS cache + tickets, and the metadata the integrator needs to overlap with captured compute. + """ + + # TransformerEngine DistributedWeight protocol (see te.pytorch.distributed_weight). + # `is_distributed_weight` is the capability marker TE dispatches on; no other + # GTP-specific state leaks to TE. + is_distributed_weight: bool = True + + # Per-chain linked-list state, keyed by chain_id; chains never cross-link (prev_w/next_w join + # only same-chain params). Call reset_gtp_state() before rebuilding a GTP model in-process. + _chain_state: Dict[str, dict] = {} + + # Recompute-forward prefetch cursor, keyed by chain_id; also cleared by reset_gtp_state(). + _recompute_chain_state: Dict[str, dict] = {} + + @classmethod + def _get_chain_state(cls, chain_id: str) -> dict: + if chain_id not in cls._chain_state: + cls._chain_state[chain_id] = { + "last_weight": None, + "link_node_count": 0, + "link_table_buffer": [], + "link_table_flushed": False, + } + return cls._chain_state[chain_id] + + @classmethod + def _get_recompute_chain_state(cls, chain_id: str) -> dict: + if chain_id not in cls._recompute_chain_state: + cls._recompute_chain_state[chain_id] = {"last_weight": None} + return cls._recompute_chain_state[chain_id] + + @classmethod + def _buffer_link_table_row( + cls, prev: "GTPShardedParam", curr: "GTPShardedParam", chain: dict + ) -> None: + """Buffer one prefetch-link row (flushed atomically on the second forward pass).""" + _W = 70 + _D = 20 + _S = 20 + + def _layer_id(name: str) -> str: + m = re.search(r"\d+", name) + return m.group() if m else "-" + + def _shape(param: "GTPShardedParam") -> str: + # Full (unsharded) weight shape that will be all-gathered across the gtp_remat + # group — i.e. the size actually prefetched into the chain, not the local shard. + try: + return str(tuple(param._unsharded_shape)) + except Exception: + return str(tuple(param.shape)) + + def _dtype(param: "GTPShardedParam") -> str: + # Report the dtype of the tensor that is ACTUALLY all-gathered, not the + # GTPShardedParam wrapper (whose logical dtype is the high-precision model-weight + # shard, i.e. params_dtype — bf16 in mixed precision). When the param has an FP8 + # representation (``param.quantized`` populated — by --fp8-param-gather's optimizer + # FP32->FP8 write, or by the per-forward cast otherwise), that quantized tensor is + # what gets gathered, yet a TE QuantizedTensor still reports a "fake" params_dtype + # ``.dtype``. So surface its raw storage dtype (e.g. uint8) tagged with the quantized + # class to make the FP8 all-gather unambiguous. + q = getattr(param, "quantized", None) + if getattr(param, "_gtp_native_fp8", False) and q is not None: + raw = getattr(q, "_rowwise_data", None) + if raw is None: + raw = getattr(q, "_data", None) + raw_dt = str(raw.dtype).replace("torch.", "") if raw is not None else "?" + return f"{type(q).__name__}/{raw_dt}" + return str(getattr(param, "dtype", "-")) + + chain["link_node_count"] += 1 + if chain["link_node_count"] == 1: + chain_id = getattr(curr, "chain_id", GTPChain.UNGRAPHED.value) + chain["link_table_buffer"].append( + f"\n[{chain_id} chain]\n{'node_id':>7} | {'layer_id':>8} |" + f" {'dtype':<{_D}} | {'shape':<{_S}} | {'curr_weight_name':<{_W}} |" + f" prev_weight_name\n{'-'*7}-+-{'-'*8}-+-{'-'*_D}-+-{'-'*_S}-+-{'-'*_W}-+-{'-'*_W}" + ) + # Seed weight (first GTP param) as row 0 + chain["link_table_buffer"].append( + f"{'0':>7} | {_layer_id(prev._debug_name):>8} | " + f"{_dtype(prev):<{_D}} | {_shape(prev):<{_S}} | {prev._debug_name:<{_W}} | -" + ) + chain["link_table_buffer"].append( + f"{chain['link_node_count']:>7} | {_layer_id(curr._debug_name):>8} | " + f"{_dtype(curr):<{_D}} | {_shape(curr):<{_S}} | " + f"{curr._debug_name:<{_W}} | {prev._debug_name}" + ) + + @staticmethod + def __new__(cls, tensor, *args, **kwargs): # pylint: disable=unused-argument + requires_grad = kwargs.get("requires_grad", True) + # pylint: disable-next=unexpected-keyword-arg + return super(GTPShardedParam, cls).__new__(cls, tensor, requires_grad=requires_grad) + + def __init__(self, tensor, *args, **kwargs): + del tensor, args, kwargs + super().__init__() + _init_gtp_runtime_attrs(self) + + @property + def _weights(self): + """Individual weight shards (self for non-routed, weight_list for routed).""" + weights = self.weight_list if self.is_routed_expert else [self] + # Only meaningful when _set_state is actively tracking transitions. + if GTP_CONFIG.check_param_states: + assert all(w.state == weights[0].state for w in weights) + return list(weights) + + @property + def _unsharded_shape_padded(self): + """Full unsharded shape *including* the pad rows on the last rank.""" + out_shape = list(self.size()) + out_shape[0] = out_shape[0] * self.group.size() + return tuple(out_shape) + + @property + def _unsharded_shape(self): + """Full unsharded shape with the pad rows stripped (logical shape).""" + out_shape = list(self._unsharded_shape_padded) + out_shape[0] -= self.pad_length + return tuple(out_shape) + + @property + def _sharded_padded_shape(self): + """This rank's local shard shape, padding included.""" + return tuple(self.size()) + + def get_padded_shard(self): + """Return the local shard already containing its share of padding (identity).""" + return self + + def _set_state(self, new_state: GTPWeightState): + """Advance the AG state (only inspected when ``check_param_states`` is on).""" + # Only inspected when check_param_states is on; skip writes otherwise. + if not GTP_CONFIG.check_param_states: + return + self.state = new_state + + def _set_rs_state(self, new_state: GTPWeightState): + """Advance the RS state (only inspected when ``check_param_states`` is on).""" + if not GTP_CONFIG.check_param_states: + return + self.rs_state = new_state + + def _double_buffer_parity(self) -> int: + """Chain-position parity (0/1) that keeps neighbouring blocks on different buffers. + + First use draws from a per-(shape, expert_idx, chain_id) counter; since first use follows + chain order, consecutive weights get 0,1,0,1... The value is cached on the param, so the + weight's fwd-AG, bwd-AG and RS buffers all share it. See ``_GTP_GROUPED_BUF_PARITY_COUNTER`` + """ + p = getattr(self, "_buf_parity", None) + if p is None: + counter_key = (self._unsharded_shape_padded, self.expert_idx, self.chain_id) + n = _GTP_GROUPED_BUF_PARITY_COUNTER.get(counter_key, 0) + p = n & 1 + _GTP_GROUPED_BUF_PARITY_COUNTER[counter_key] = n + 1 + self._buf_parity = p + return p + + def _get_cache_key(self, dtype, fwd: bool, reduce_scatter: bool) -> tuple: + """Build cache key from output shape + dtype. + + Weights with matching gathered shape and dtype share a buffer. For experts gathered + in parallel, self.expert_idx keeps each distinct; same-indexed experts across layers share. + + Grouped one-block-ahead chains additionally fold in a double-buffer parity so a prefetched + layer N+1 weight never lands in the buffer that layer N is still consuming (see + ``_GTP_GROUPED_BUF_PARITY_COUNTER``). + """ + + if not isinstance(dtype, torch.dtype): + key = ( + self._unsharded_shape_padded, + dtype, + fwd, + not fwd, + self.expert_idx, + reduce_scatter, + ) + else: + key = (self._unsharded_shape_padded, dtype, self.expert_idx, reduce_scatter) + if _chain_is_grouped(self.chain_id): + # chain_id keeps fc1/fc2 apart (both can be in flight at once, even if same-shaped); + # parity alternates consecutive blocks between two buffers. + key = key + (self.chain_id, self._double_buffer_parity()) + return key + + def _strip_padding(self, tensor): + if self.pad_length == 0: + return tensor + + if isinstance(tensor, QuantizedTensor): + assert isinstance( + tensor, (NVFP4TensorStorage, MXFP8TensorStorage) + ), f"Unsupported quantized tensor type for GTP padding: {type(tensor)}" + + metadata = tensor.get_metadata() + if metadata.get("rowwise_data") is not None: + metadata["rowwise_data"] = metadata["rowwise_data"][: -self.pad_length] + if metadata.get("columnwise_data") is not None: + if isinstance(tensor, NVFP4TensorStorage): + # NVFP4 transposes columnwise and packs 2 values per byte + metadata["columnwise_data"] = metadata["columnwise_data"][ + ..., : -self.pad_length // 2 + ].contiguous() + else: + # MXFP8 columnwise is not transposed, strip first dim + metadata["columnwise_data"] = metadata["columnwise_data"][: -self.pad_length] + M = self._unsharded_shape[0] + if isinstance(tensor, NVFP4TensorStorage): + # NVFP4 scale_inv shapes (see NVFP4Quantizer.get_scale_shape): + # rowwise_scale_inv: [round_up(M, 128), round_up(ceil(K/16), 4)] + # columnwise_scale_inv: [round_up(K, 128), round_up(ceil(M/16), 4)] + # GTP shards M (dim 0 of the weight), so strip to the unpadded sizes. + if metadata.get("rowwise_scale_inv") is not None: + m_rows = round_up_to_nearest_multiple(M, 128) + metadata["rowwise_scale_inv"] = metadata["rowwise_scale_inv"][:m_rows] + if metadata.get("columnwise_scale_inv") is not None: + m_tiles = round_up_to_nearest_multiple( + math.ceil(M / NVFP4_BLOCK_SCALING_SIZE), 4 + ) + metadata["columnwise_scale_inv"] = metadata["columnwise_scale_inv"][ + :, :m_tiles + ].contiguous() + else: + # MXFP8 scale_inv shapes (see MXFP8Quantizer.get_scale_shape): + # rowwise_scale_inv: [round_up(M, 128), round_up(K//32, 4)] + # columnwise_scale_inv: [round_up(M//32, 4), round_up(K, 128)] + # GTP shards M (dim 0 of the weight), so strip to the unpadded sizes. + if metadata.get("rowwise_scale_inv") is not None: + m_rows = round_up_to_nearest_multiple(M, 128) + metadata["rowwise_scale_inv"] = metadata["rowwise_scale_inv"][:m_rows] + if metadata.get("columnwise_scale_inv") is not None: + m_tiles = round_up_to_nearest_multiple(M // MXFP8_BLOCK_SCALING_SIZE, 4) + metadata["columnwise_scale_inv"] = metadata["columnwise_scale_inv"][:m_tiles] + + return type(tensor)(**metadata, shape=self._unsharded_shape, dtype=torch.bfloat16) + + return tensor[: -self.pad_length] + + def _all_gather_weight(self, async_op, fwd, nvtx_label=None): + """Quantize (if needed) and all-gather weight. Returns (weight_total, handle).""" + if nvtx_label is None: + nvtx_label = ( + self._debug_name + (".fwd" if fwd else ".bwd") + (".async" if async_op else ".sync") + ) + nvtx_range_push(f"{nvtx_label}.all_gather_weight") + + weights = self._weights + + # 1. Transition state for async gathers. Skip during recompute-forward: it gathers + # rowwise (_ag_ticket_fwd) while a bwd-chain prefetch may hold an in-flight columnwise + # AG state (_ag_ticket_bwd) on the same weight — clobbering breaks the dgrad consume. + if GTP_CONFIG.check_param_states and not in_fp8_activation_recompute_phase(): + new_state = GTPWeightState.ASYNC_WAIT if async_op else GTPWeightState.DATA_READY_SYNC + for w in weights: + w._set_state(new_state) + + # 2. Set FP8 usage direction (rowwise for fwd, columnwise for bwd) on the GATHER + # quantizer copy — NEVER on the param's own quantizer: the optimizer's + # copy_ -> quantize_ update writes whichever usages the param quantizer has enabled, + # so leaving it rowwise=False after a bwd gather would freeze the rowwise (fwd) data. + # No re-quantize here: with mxfp8 + --fp8-param-gather the shard already IS a native + # FP8 tensor. BF16 GTP carries no quantizer and gathers the BF16 shard as-is. + native_fp8 = getattr(self, "_gtp_native_fp8", False) + quantizers = [getattr(w, "_gtp_gather_quantizer", None) for w in weights] + if native_fp8: + for q in quantizers: + q.set_usage(rowwise=fwd, columnwise=not fwd) + + # 3. Build gather inputs. The gather collective takes the per-weight quantizers so it can + # reconstruct the gathered FP8 tensor's scale/metadata (None for BF16 GTP). + if native_fp8: + gather_weights = [w.quantized for w in weights] + else: + gather_weights = list(w.get_padded_shard() for w in weights) + + # 4. Cache checkout — use pooled buffers for both async and sync gathers + # to avoid allocating fresh memory each iteration. gather-buffer dtypes are stable + # post-construction (FP8 quantizer dtype for native-FP8 shards, else the BF16 dtype), + # so cache them on the anchor (self == weights[0]) instead of rebuilding each call. + dtypes = self._cached_dtypes + if dtypes is None: + dtypes = [q.dtype if q is not None else w.dtype for q, w in zip(quantizers, weights)] + self._cached_dtypes = dtypes + out_buffers = [] + cache = get_global_GTP_cache() + for p, dt in zip(weights, dtypes): + if fwd: + if p._ag_ticket_fwd is None: + p._ag_ticket_fwd = cache.reserve(p, dt, fwd=True) + cache.get(p._ag_ticket_fwd) + cache.release(p._ag_ticket_fwd) + out_buffers.append(cache.get(p._ag_ticket_fwd)) + else: + if p._ag_ticket_bwd is None: + p._ag_ticket_bwd = cache.reserve(p, dt, fwd=False) + out_buffers.append(cache.get(p._ag_ticket_bwd)) + + # 5. Communicate. + gtp_remat_group = self._cached_gtp_remat_group + if gtp_remat_group is None: + gtp_remat_group = weights[0].group + self._cached_gtp_remat_group = gtp_remat_group + if GTP_CONFIG.check_param_states and len(gather_weights) > 1: + # Debug invariant: batched AG needs distinct output buffers per expert. + assert len(set(id(b) for b in out_buffers)) == len( + out_buffers + ), "Duplicate output buffers in batched all-gather — experts need distinct cache keys" + + # ASYNC AG: issue on ag_stream so its tail reflects the collective's full lifecycle + # (what external wait_stream(ag_stream) drains depend on). The explicit outer→ag_stream + # sync event preserves the upstream quantize-writer edge the bare stream context drops; + # held on self so the event pool can't recycle it between capture and replay. + # SYNC AG: stay on caller — output ready on return. + if async_op: + outer_stream = torch.cuda.current_stream() + ag_stream = get_ag_stream(self.chain_id, gtp_remat_group) + if getattr(self, "_ag_outer_sync_event", None) is None: + self._ag_outer_sync_event = torch.cuda.Event() + outer_sync_event = self._ag_outer_sync_event + outer_sync_event.record(outer_stream) + ag_stream.wait_event(outer_sync_event) + ag_ctx = torch.cuda.stream(ag_stream) + else: + ag_ctx = nullcontext() + + with ag_ctx: + if len(gather_weights) > 1: + nvtx_range_push(f"{nvtx_label}.batched_gtp_ag") + results, handle = grouped_gather_along_first_dim( + gather_weights, + gtp_remat_group, + async_op=async_op, + quantizers=quantizers, + output_tensors=out_buffers, + ) + nvtx_range_pop(f"{nvtx_label}.batched_gtp_ag") + else: + nvtx_range_push(f"{nvtx_label}.gtp_ag") + weight_total, handle = gather_along_first_dim( + gather_weights[0], + gtp_remat_group, + quantizer=quantizers[0], + async_op=async_op, + output_tensor=out_buffers[0] if out_buffers is not None else None, + ) + nvtx_range_pop(f"{nvtx_label}.gtp_ag") + results = [weight_total] + + result = results if self.is_routed_expert else results[0] + + # 6. Wrap handle. + if async_op: + handle = GTPShardHandle(handle, weights) + else: + handle = None + + nvtx_range_pop(f"{nvtx_label}.all_gather_weight") + return result, handle + + def _wait_param_gather(self): + # Enter ag_stream context so handle.wait() + ag_event.record() both + # land on ag_stream. That makes ag_event mark ag_stream's tail, which + # is what external drains via wait_stream(ag_stream) actually block on. + ag_stream = self._cached_ag_stream + if ag_stream is None: + ag_stream = get_ag_stream(self.chain_id, self.group) + self._cached_ag_stream = ag_stream + with torch.cuda.stream(ag_stream): + if self._prefetch_handle is not None: + self._prefetch_handle.wait() + self._prefetch_handle = None + self.ag_event.record() + + def _all_gather_weight_on_demand(self, fwd): + result, _ = self._all_gather_weight(async_op=False, fwd=fwd) + result = result if self.is_routed_expert else [result] + result = [self._strip_padding(r) for r in result] + result = [r.detach().requires_grad_(w.requires_grad) for r, w in zip(result, self._weights)] + return result if self.is_routed_expert else result[0] + + def _get_prefetched_weight(self, fwd): + # Stale-read guard: state must reflect an AG issued for this cycle; + # otherwise cache.get() would return the prior iter's AG buffer. + if GTP_CONFIG.check_param_states: + for w in self._weights: + assert w.state in ( + GTPWeightState.ASYNC_WAIT, + GTPWeightState.DATA_READY, + GTPWeightState.DATA_READY_SYNC, + ), ( + f"[GTP] _get_prefetched_weight({'fwd' if fwd else 'bwd'}) on " + f"{self._debug_name} with state={w.state!r} — no AG issued; " + "cache.get() would return stale data. Check the chain's " + "_need_weight_prefetch flag and issuer's prefetch logic." + ) + _was_drained = getattr(self, "_already_ag_drained", False) + if _was_drained: + # Producer already drained via wait_async_comms; skip the captured cross-graph + # wait (a CUDA no-op anyway). Correctness comes from the eager main_stream sync. + self._already_ag_drained = False + else: + # Intra-graph or eager consume: drain inline. + self._wait_param_gather() + self.ag_event.wait() + + # Retrieve prefetched results from cache + result = [] + cache = get_global_GTP_cache() + for w in self._weights: + ticket = w._ag_ticket_fwd if fwd else w._ag_ticket_bwd + result.append(cache.get(ticket)) + + result = [self._strip_padding(r) for r in result] + + result = [r.detach().requires_grad_(w.requires_grad) for r, w in zip(result, self._weights)] + return result if self.is_routed_expert else result[0] + + def _wait_recompute_param_gather(self): + # Recompute-chain analogue of _wait_param_gather, on the _recompute_* slot. + ag_stream = self._cached_ag_stream + if ag_stream is None: + ag_stream = get_ag_stream(self.chain_id, self.group) + self._cached_ag_stream = ag_stream + with torch.cuda.stream(ag_stream): + if self._recompute_prefetch_handle is not None: + self._recompute_prefetch_handle.wait() + self._recompute_prefetch_handle = None + self._recompute_ag_event.record() + + def _recompute_prefetch_next(self, target, nvtx_label=None): + # Issue target's rowwise (fwd) AG into its recompute slot. _all_gather_weight skips the + # AG-state transition under recompute, so target's dgrad state is untouched; result lands + # in target._ag_ticket_fwd. + _, handle = target._all_gather_weight(async_op=True, fwd=True, nvtx_label=nvtx_label) + target._recompute_prefetch_handle = handle + + def _get_recompute_prefetched_weight(self): + # Recompute-chain analogue of _get_prefetched_weight (state-neutral; reads the + # rowwise _ag_ticket_fwd via the _recompute_* slot). + if self._recompute_already_drained: + # Producer already drained via wait_async_comms (CG capture); skip the + # captured cross-graph wait (CUDA no-op anyway). + self._recompute_already_drained = False + else: + self._wait_recompute_param_gather() + self._recompute_ag_event.wait() + + result = [] + cache = get_global_GTP_cache() + for w in self._weights: + result.append(cache.get(w._ag_ticket_fwd)) + result = [self._strip_padding(r) for r in result] + result = [r.detach().requires_grad_(w.requires_grad) for r, w in zip(result, self._weights)] + return result if self.is_routed_expert else result[0] + + def all_gather_and_prefetch_bwd(self, nvtx_label=None): + """Backward variant: get the current weight (cached if prefetched, else sync gather) + and async-prefetch prev_w. + + Safe via the coat-check cache: get() returns the current buffer to the pool, and the + prefetch's checkout allocates a separate buffer if the pool is empty (current buffer + still live via the caller's reference). + + Returns: + weight_total + """ + + if GTP_CONFIG.weight_prefetch and self.next_w is not None: + result = self._get_prefetched_weight(False) + else: + result = self._all_gather_weight_on_demand(False) + + if ( + GTP_CONFIG.weight_prefetch + and self.prev_w is not None + and self.prev_w._need_weight_prefetch + and self.prev_w._need_weight_prefetch_bwd + ): + # Pre-AG work (quantize, ticket lookup) runs on caller's stream; the NCCL collective + # is wrapped on ag_stream inside _all_gather_weight (see its async/sync gate). + _, handle = self.prev_w._all_gather_weight( + async_op=True, fwd=False, nvtx_label=nvtx_label + ) + self.prev_w._prefetch_handle = handle + + # The unsharded tensor has been returned, no pending work so reset state to NONE + if GTP_CONFIG.check_param_states: + for w in self._weights: + w._set_state(GTPWeightState.NONE) + + if GTP_CONFIG.weight_prefetch and self.next_w is not None: + cache = get_global_GTP_cache() + for w in self._weights: + cache.release(w._ag_ticket_bwd) + + return result + + def batched_all_gather_and_prefetch_bwd(self, nvtx_label=None): + """Batched backward all-gather + prefetch. Wrapper around all_gather_and_prefetch_bwd.""" + assert self.is_routed_expert and self.weight_list is not None + return self.all_gather_and_prefetch_bwd(nvtx_label=nvtx_label) + + def all_gather_and_prefetch(self, fwd: bool = True, nvtx_label: str = None): + """All-gather the current weight and async-prefetch the next. + + Returns: + weight_total + """ + # During an activation-recompute forward (runs in backward), route consume + + # prefetch through the recompute-forward chain on its own _recompute_* slot + # (see __init__) instead of the fwd/bwd chains; lazy-built below. + in_recompute = in_fp8_activation_recompute_phase() + use_recompute_chain = in_recompute and GTP_CONFIG.weight_prefetch + + # Consume current weight. + if use_recompute_chain and self._recompute_prev is not None: + result = self._get_recompute_prefetched_weight() + elif not in_recompute and GTP_CONFIG.weight_prefetch and self.prev_w is not None: + result = self._get_prefetched_weight(True) + else: + # On-demand: chain head (fwd or recompute global-first) or first-iter build. + result = self._all_gather_weight_on_demand(True) + + # Prefetch next weight on the matching chain. + if ( + use_recompute_chain + and self._recompute_next is not None + and self._recompute_next._need_weight_prefetch + ): + self._recompute_prefetch_next(self._recompute_next, nvtx_label=nvtx_label) + elif ( + not in_recompute + and GTP_CONFIG.weight_prefetch + and self.next_w is not None + and self.next_w._need_weight_prefetch + ): + # Pre-AG work on caller; NCCL wrap lives at the collective site + # inside _all_gather_weight. See all_gather_and_prefetch_bwd. + _, handle = self.next_w._all_gather_weight( + async_op=True, fwd=fwd, nvtx_label=nvtx_label + ) + self.next_w._prefetch_handle = handle + + # Unsharded tensor returned, no pending work → reset state to NONE. Skip during recompute: + # a bwd-chain prefetch may hold an in-flight AG state this weight's later dgrad needs. + if GTP_CONFIG.check_param_states and not in_recompute: + for w in self._weights: + w._set_state(GTPWeightState.NONE) + + cls = type(self) + + # Lazy-build the recompute-forward prefetch chain (first backward, in recompute order). + # Consume/prefetch above used the prior iter's links, so the first backward runs on-demand + # while these are established. + if in_recompute and not self._recompute_initialized: + rchain = cls._get_recompute_chain_state(self.chain_id) + last_r = rchain["last_weight"] + if last_r is not None and last_r._recompute_next is None: + last_r._recompute_next = self + self._recompute_prev = last_r + self._recompute_initialized = True + rchain["last_weight"] = self + + # Lazy population of the fwd/bwd linked list: link previous weight to current. + # Uses per-chain state so dense and expert chains never cross-link. + chain = cls._get_chain_state(self.chain_id) + if not self.prefetch_initialized: + last_w = chain["last_weight"] + if last_w is not None and last_w.next_w is None: + cls._buffer_link_table_row(last_w, self, chain) + last_w.next_w = self + self.prev_w = last_w + + cache = get_global_GTP_cache() + + # Set the fwd ag buffer (gather quantizer copy — the param's own quantizer is + # reserved for the optimizer's update path; see attach_gtp_to_presharded_module). + quantizers = [getattr(w, "_gtp_gather_quantizer", None) for w in self._weights] + dtypes = [ + q.dtype if q is not None else w.dtype for q, w in zip(quantizers, self._weights) + ] + for w, dt in zip(self._weights, dtypes): + w._ag_ticket_fwd = cache.reserve(w, dt, fwd=True) + cache.get(w._ag_ticket_fwd) + cache.release(w._ag_ticket_fwd) + + self.prefetch_initialized = True + chain["last_weight"] = self + elif not chain["link_table_flushed"] and chain["link_table_buffer"]: + # Second forward pass: flush the complete table atomically to avoid interleaving + chain["link_table_flushed"] = True + log_single_rank(logger, logging.INFO, "\n".join(chain["link_table_buffer"]) + "\n") + + return result + + def batched_all_gather_and_prefetch(self, **kwargs): + """Batched all-gather + prefetch for expert weights (wraps all_gather_and_prefetch).""" + assert self.is_routed_expert and self.weight_list is not None + return self.all_gather_and_prefetch(**kwargs) + + def get_wgrad_tensor(self): + """Pool-allocate a wgrad scratch tensor of unsharded shape for the bwd GEMM.""" + return _wgrad_pool_get(self._unsharded_shape, self.main_grad.dtype, self.device) + + def register_grad_accum_hook(self, grad_accum_node, hook): + """Register a DDP backward hook to call after the wgrad RS finalize. + + For GTP params autograd may receive None (async RS), so the normal grad-accumulator + hook never fires; the integrator (Graphed.backward for captured chains, or the eager + chain-tail cascade) calls this hook explicitly after RS wait + accumulation, so DDP's + register_grad_ready fires at the right time. grad_accum_node is accepted for API + compatibility but not retained — only the hook callable. + """ + del grad_accum_node + self._grad_accum_hook = hook + + @staticmethod + def _handle_megatron_grad_accum(param): + """Handle megatron DDP and gradient-accumulation fusion. + + Do NOT set param.grad before calling the hook — the hook checks param.grad and would + accumulate it into main_grad if zero_out_wgrad is True, corrupting it with a dummy. + + Returns a cached dummy wgrad; sync callers use it as the graph-safe grad, async drains + discard it. + """ + if hasattr(param, "grad_added_to_main_grad"): + param.grad_added_to_main_grad = True + dummy_grad = get_dummy_wgrad(list(param.main_grad.shape), param.dtype) + if getattr(param, "_grad_accum_hook", None) is not None: + param._grad_accum_hook() + + param._set_rs_state(GTPWeightState.NONE) + return dummy_grad + + def _wait_reduce_scatter(self, finalize_grad=False): + # Enter rs_stream context so handle.wait() + rs_event.record() land on rs_stream + # (mirrors _wait_param_gather). With finalize_grad=True, main_grad.add_ also runs on + # rs_stream right after the NCCL RS — starts during AG drain, not after, avoiding + # SM-saturation that blocks cross-graph overlap. + rs_stream = self._cached_rs_stream + if rs_stream is None: + rs_stream = get_rs_stream(self.chain_id, self.group) + self._cached_rs_stream = rs_stream + with torch.cuda.stream(rs_stream): + if self._wgrad_rs_handle is not None: + self._wgrad_rs_handle.wait() + self._wgrad_rs_handle = None + self.rs_event.record() + if finalize_grad: + cache = get_global_GTP_cache() + for w in self._weights: + wgrad_rs = cache.get(w._rs_ticket) + w.main_grad.add_(wgrad_rs) + cache.release(w._rs_ticket) + # Fire grad-ready AFTER all adds (separate loop so a bucket-completing + # grad-ready can't dispatch the RS before a sibling's add). With autograd + # grad-ready suppressed for GTP params (DDP register_grad_accum_hook), this + # is the only grad-ready for a weight finalized here; else the bucket orphans. + for w in self._weights: + self._handle_megatron_grad_accum(w) + self._already_finalized = True + # Release stashed wgrad inputs: UNGRAPHED buffers go back to the pool; + # GRAPHED just drops Python refs (addresses must stay stable for CG). + if getattr(self, "_wgrad_input_bufs", None) is not None: + if not _chain_is_graphed(self.chain_id): + for buf in self._wgrad_input_bufs: + _wgrad_pool_put(buf) + self._wgrad_input_bufs = None + + def _prescale_wgrads_for_mean_rs(self, wgrads): + """Pre-scale wgrad by 1/gtp_remat so the SUM reduce-scatter yields the gtp_remat mean. + + Single choke point for every RS path. Composes with DDP's 1/replicate prescale and + finalize's AVG to give the full (replicate x gtp_remat) mean. Skipped under + calculate_per_token_loss, where DDP does no 1/dp scaling and total_global_tokens (which + counts gtp_remat peers' tokens) normalizes instead — there the gtp_remat axis must SUM + like the DP axis (a 1/gtp_remat mean would shrink every gtp_remat grad). + """ + gtp_remat_size = self.group.size() + if gtp_remat_size > 1 and not GTP_CONFIG.calculate_per_token_loss: + torch._foreach_mul_(list(wgrads), 1.0 / gtp_remat_size) + + def _reduce_scatter(self, wgrads, async_op, nvtx_label=None): + """Reduce-scatter one or more wgrads → (outputs, handle). Single tensor: plain RS; + multiple: coalesced RS.""" + if nvtx_label is None: + nvtx_label = self._debug_name + ".bwd" + (".async" if async_op else ".sync") + + # MEAN reduce-scatter: pre-scale wgrad so the SUM collective yields the gtp_remat mean. + self._prescale_wgrads_for_mean_rs(wgrads) + + if GTP_CONFIG.check_param_states: + new_rs_state = GTPWeightState.ASYNC_WAIT if async_op else GTPWeightState.DATA_READY_SYNC + for w in self._weights: + w._set_rs_state(new_rs_state) + + if self.pad_length > 0: + wgrads = [torch.nn.functional.pad(w, (0, 0, 0, self.pad_length)) for w in wgrads] + + if async_op: + dtypes = [w.dtype for w in wgrads] + out_buffers = [] + cache = get_global_GTP_cache() + for p, dt in zip(self._weights, dtypes): + if p._rs_ticket is None: + p._rs_ticket = cache.reserve(p, dt, fwd=False, reduce_scatter=True) + out_buffers.append(cache.get(p._rs_ticket)) + else: + out_buffers = [None] * len(wgrads) + + # ASYNC RS: issue on rs_stream so its tail reflects the collective's full lifecycle + # (what external wait_stream(rs_stream) drains depend on). The explicit outer→rs_stream + # sync event preserves the wgrad-GEMM writer edge the bare stream context drops; held on + # self so the event pool can't recycle it between capture and replay. Mirrors the AG path. + # SYNC RS: stay on caller — output ready on return. + if async_op: + outer_stream = torch.cuda.current_stream() + rs_stream = get_rs_stream(self.chain_id, self.group) + if getattr(self, "_rs_outer_sync_event", None) is None: + self._rs_outer_sync_event = torch.cuda.Event() + outer_sync_event = self._rs_outer_sync_event + outer_sync_event.record(outer_stream) + rs_stream.wait_event(outer_sync_event) + rs_ctx = torch.cuda.stream(rs_stream) + else: + rs_ctx = nullcontext() + + with rs_ctx: + if len(wgrads) == 1: + nvtx_range_push(f"{nvtx_label}.gtp_rs") + out, handle = reduce_scatter_along_first_dim( + wgrads[0], self.group, async_op=async_op, output=out_buffers[0] + ) + nvtx_range_pop(f"{nvtx_label}.gtp_rs") + return [out], handle + + outputs = [] + nvtx_range_push(f"{nvtx_label}.batched_gtp_rs") + with torch.distributed._coalescing_manager( + group=self.group, device=wgrads[0].device, async_ops=async_op + ) as cm: + for out_buffer, tensor in zip(out_buffers, wgrads): + out, _ = reduce_scatter_along_first_dim(tensor, self.group, output=out_buffer) + outputs.append(out) + nvtx_range_pop(f"{nvtx_label}.batched_gtp_rs") + + return outputs, cm if async_op else None + + def wgrad_reduce_scatter(self, wgrad, nvtx_label=None): + """Reduce-scatter wgrad(s): sync for the last weight, async+deferred for others. + Accepts a single tensor (non-routed) or a list (routed experts). + + Returns: + Single tensor or list for sync (last weight) — backward returns this. + None or tuple of Nones for async — backward returns this. + """ + batched = isinstance(wgrad, (list, tuple)) + wgrads = list(wgrad) if batched else [wgrad] + weights = self._weights + + # UNGRAPHED wgrads recycle via the standalone pool (_wgrad_pool_put); GRAPHED wgrads + # cannot, since CUDA graphs require stable buffer addresses across replay. + poolable = not _chain_is_graphed(self.chain_id) + + if GTP_CONFIG.async_reduction and self.prev_w is not None: + # Async RS (not last weight — deferred finish). Pre-RS work on caller; NCCL wrap + # lives at the collective site inside _reduce_scatter (mirrors the AG prefetch sites). + _, rs_handle = self._reduce_scatter(wgrads, async_op=True, nvtx_label=nvtx_label) + self._wgrad_rs_handle = GTPShardHandle(rs_handle, weights, reduce_scatter=True) + # Stash wgrad input buffers — cannot recycle yet because the async RS + # kernel is still reading them on rs_stream. + self._wgrad_input_bufs = wgrads + ret = tuple([None] * len(wgrads)) if batched else None + else: + # Sync reduce-scatter — reached as the natural chain-head case, recycle immediately + wgrads, _ = self._reduce_scatter(wgrads, async_op=False, nvtx_label=nvtx_label) + nvtx_range_push(f"{nvtx_label}.gtp_wgrad_accum") + if len(weights) == 1: + weights[0].main_grad.add_(wgrads[0]) + else: + torch._foreach_add_([p.main_grad for p in weights], wgrads) + nvtx_range_pop(f"{nvtx_label}.gtp_wgrad_accum") + result = [self._handle_megatron_grad_accum(p) for p in weights] + + if poolable: + for buf in wgrads: + _wgrad_pool_put(buf) + ret = result if batched else result[0] + + # Wait for last reduce scatter if it was async + # Currently only support reduce scattering in reverse order + if GTP_CONFIG.async_reduction and self.next_w is not None: + self.next_w._wait_reduce_scatter() + + if getattr(self.next_w, "_already_finalized", False): + self.next_w._already_finalized = False + else: + self.next_w.rs_event.wait() + cache = get_global_GTP_cache() + next_weights = self.next_w._weights + wgrads = [cache.get(w._rs_ticket) for w in next_weights] + nvtx_range_push(f"{self.next_w._debug_name}.gtp_wgrad_accum_deferred") + # Only batch with _foreach_add_ when finalizing multiple (routed) weights. + if len(next_weights) == 1: + next_weights[0].main_grad.add_(wgrads[0]) + else: + torch._foreach_add_([w.main_grad for w in next_weights], wgrads) + nvtx_range_pop(f"{self.next_w._debug_name}.gtp_wgrad_accum_deferred") + for w in next_weights: + self._handle_megatron_grad_accum(w) + cache.release(w._rs_ticket) + + return ret + + def batched_wgrad_reduce_scatter(self, wgrad_list, nvtx_label=None): + """Batched version of wgrad_reduce_scatter.""" + assert self.is_routed_expert and self.weight_list is not None + return self.wgrad_reduce_scatter(wgrad_list, nvtx_label=nvtx_label) + + # ------------------------------------------------------------------ + # TransformerEngine DistributedWeight protocol. TE's fwd/bwd dispatch through these generic + # names (see te.pytorch.distributed_weight.materialize_weights_for_forward et al.). The leader + # param encapsulates the whole group via self._weights, so a single call covers both the Linear + # (one weight) and GroupedLinear (routed-expert list) cases; the underlying methods already + # return a single tensor or a list accordingly. These are thin adapters over the GTP methods. + # ------------------------------------------------------------------ + def materialize_group_for_forward(self): + """Protocol: all-gather the group's shard(s) for the forward GEMM.""" + return self.all_gather_and_prefetch(fwd=True) + + def materialize_group_for_backward(self, nvtx_label=None): + """Protocol: re-materialize the group's weight(s) for the backward GEMMs.""" + return self.all_gather_and_prefetch_bwd(nvtx_label=nvtx_label) + + def finalize_group_grads(self, wgrads, nvtx_label=None): + """Protocol: reduce-scatter the group's freshly computed weight grad(s).""" + return self.wgrad_reduce_scatter(wgrads, nvtx_label=nvtx_label) + + def grad_buffer(self): + """Protocol: the wgrad accumulation scratch buffer for this weight.""" + return self.get_wgrad_tensor() + + def get_data_tensors(self): + """Expose self as the lone data tensor for TE's offload-marking interface. + + TE's mark_activation_offload treats any non-plain tensor as a storage wrapper and calls + get_data_tensors() on it; a sharded param has no inner buffers, so it is its own. + """ + return (self,) + + def __torch_function__(self, func, types, args=(), kwargs=None): + """Subclass-preserving dispatch for ``detach`` (other ops fall through).""" + del types # required by protocol, unused here + if kwargs is None: + kwargs = {} + + if func is torch.Tensor.detach: + with torch._C.DisableTorchFunctionSubclass(): + # Perform the raw detach + result = func(*args, **kwargs) + # Re-wrap it in your subclass so PyTorch is happy + return result.as_subclass(type(self)) + + # 2. For everything else (add, mul, etc.), be transparent/decay. + with torch._C.DisableTorchFunctionSubclass(): + return func(*args, **kwargs) + + +@dataclass +class _TicketSlot: + """Internal slot backing a persistent ticket in the GTP buffer cache.""" + + key: tuple # cache key (shape, dtype, ...) + param: "GTPShardedParam" # for lazy allocation metadata + dtype: object # torch.dtype or tex.DType + reduce_scatter: bool + fwd: bool + chain_id: str = GTPChain.GRAPHED.value # chain this slot belongs to + buf: Optional[torch.Tensor] = field(default=None) # None when released or after clear() + + +# CUDA-graph memory pool: routes GRAPHED-chain allocations (AG/RS buffers, quantized weight +# storage) into the capture pool at creation time, avoiding post-hoc reallocation. Registered +# via set_cuda_graph_mempool before the first graphed forward; stays None when CG is off, where +# _graphed_alloc is a no-op (regular allocator). +_CG_MEMPOOL_DEVICE = None +_CG_MEMPOOL = None + + +def set_cuda_graph_mempool(device, mempool): + """Register the CUDA-graph memory pool for GRAPHED-chain GTP allocations.""" + global _CG_MEMPOOL_DEVICE, _CG_MEMPOOL + _CG_MEMPOOL_DEVICE = device + _CG_MEMPOOL = mempool + + +@contextmanager +def _graphed_alloc(chain_id): + """Route allocations in this block into the registered CG mempool when ``chain_id`` + is GRAPHED and a pool is registered; otherwise a no-op (regular allocator).""" + if _CG_MEMPOOL is not None and _chain_is_graphed(chain_id): + torch._C._cuda_beginAllocateCurrentThreadToPool(_CG_MEMPOOL_DEVICE, _CG_MEMPOOL) + try: + yield + finally: + torch._C._cuda_endAllocateToPool(_CG_MEMPOOL_DEVICE, _CG_MEMPOOL) + else: + yield + + +class GTPWeightCache: + """Ticket-based buffer pool for GTP all-gather / reduce-scatter buffers. + + - reserve(param, dtype, fwd) → ticket: assign a persistent ticket (no buffer yet). + - get(ticket) → buffer: return the buffer, lazily (re)allocating from pool or fresh. + - release(ticket): return the buffer to the pool; ticket stays valid. + - clear(): drop all buffers/pools; tickets stay valid, next get() allocates fresh. + """ + + # Bytes per element for known dtypes (for logging). Add entries when GTP caches buffers of + # new quantized dtypes — only DType values the TE pybind bindings expose (verify via + # hasattr(tex.DType, ...) before adding speculative entries). + _BYTES_PER_ELEMENT = { + torch.bfloat16: 2, + torch.float16: 2, + torch.float32: 4, + tex.DType.kFloat4E2M1: 0.5, + tex.DType.kFloat8E4M3: 1, + tex.DType.kFloat8E5M2: 1, + } + + def __init__(self): + self._pool: Dict[tuple, List[torch.Tensor]] = defaultdict(list) + self._slots: Dict[int, _TicketSlot] = {} + self._next_ticket: int = 0 + self._total_bytes: int = 0 # running total of allocated bytes + self.key_to_allocate_func = {} + + @staticmethod + def _buf_bytes(shape, dtype) -> int: + """Estimate buffer size in bytes.""" + numel = 1 + for d in shape: + numel *= d + if dtype not in GTPWeightCache._BYTES_PER_ELEMENT: + raise KeyError( + f"GTPWeightCache._buf_bytes: unknown dtype {dtype!r}. " + "Add it to GTPWeightCache._BYTES_PER_ELEMENT with its bytes-per-element." + ) + return int(numel * GTPWeightCache._BYTES_PER_ELEMENT[dtype]) + + def _allocate_buffer( + self, param: "GTPShardedParam", dtype, reduce_scatter, fwd + ) -> torch.Tensor: + if reduce_scatter: + out_shape = param._sharded_padded_shape + else: + out_shape = param._unsharded_shape_padded + + # Route GRAPHED-chain buffers into the CG mempool at creation (see _graphed_alloc). + with _graphed_alloc(getattr(param, "chain_id", GTPChain.UNGRAPHED.value)): + if not isinstance(dtype, torch.dtype): + # Use the gather quantizer copy: mutating the param's own quantizer usage + # would corrupt the optimizer's quantize_ update direction (frozen weights). + quantizer = getattr(param, "_gtp_gather_quantizer", None) or param._quantizer + assert quantizer is not None + quantizer.set_usage(rowwise=fwd, columnwise=not fwd) + + buf = quantizer.make_empty( + out_shape, dtype=torch.bfloat16, device=torch.cuda.current_device() + ) + else: + buf = torch.empty( + out_shape, + dtype=dtype, + device=param.device, + memory_format=torch.contiguous_format, + ) + + buf_bytes = self._buf_bytes(out_shape, dtype) + self._total_bytes += buf_bytes + dtype_str = ( + str(dtype) if isinstance(dtype, torch.dtype) else getattr(dtype, "name", str(dtype)) + ) + op_str = "RS(grad)" if reduce_scatter else ("AG(fwd)" if fwd else "AG(bwd)") + log_single_rank( + logger, + logging.INFO, + f"[GTP Cache] +{buf_bytes / 1024**2:.1f} MB (shape={out_shape}, dtype={dtype_str}) " + f"total={self._total_bytes / 1024**2:.1f} MB param: {param._debug_name} " + f"op: {op_str}", + ) + return buf + + def reserve(self, param: "GTPShardedParam", dtype, fwd: bool, reduce_scatter=False) -> int: + """Assign a persistent ticket. No buffer is allocated until ``get()``.""" + key = param._get_cache_key(dtype, fwd, reduce_scatter) + ticket = self._next_ticket + self._next_ticket += 1 + + self._slots[ticket] = _TicketSlot( + key=key, + param=param, + dtype=dtype, + reduce_scatter=reduce_scatter, + fwd=fwd, + chain_id=getattr(param, "chain_id", GTPChain.UNGRAPHED.value), + ) + return ticket + + def get(self, ticket: int) -> torch.Tensor: + """Return the buffer for *ticket*, lazily allocating if needed.""" + slot = self._slots[ticket] + if slot.buf is None: + pool = self._pool[slot.key] + slot.buf = ( + pool.pop() + if pool + else self._allocate_buffer( + slot.param, slot.dtype, slot.reduce_scatter, fwd=slot.fwd + ) + ) + self.key_to_allocate_func[slot.key] = ( + slot.param, + slot.dtype, + slot.reduce_scatter, + slot.fwd, + ) + + return slot.buf + + def release(self, ticket: int): + """Return the buffer to the pool (ticket stays valid). + + slot.buf is intentionally NOT cleared: get() must stay idempotent so CUDA-graph-captured + buffers keep their fixed address across replays. + """ + slot = self._slots[ticket] + if slot.buf is None: + return + # Use identity check — tensor == tensor returns a multi-element bool tensor + # which crashes in a boolean context ("Boolean value of Tensor is ambiguous"). + if not any(b is slot.buf for b in self._pool.get(slot.key, [])): + self._pool[slot.key].append(slot.buf) + + def clear(self): + """Drop all buffers; tickets remain valid and lazily re-allocate on next get().""" + for slot in self._slots.values(): + slot.buf = None + self._pool.clear() + self._total_bytes = 0 + + +def get_global_GTP_cache() -> GTPWeightCache: + """Get or lazily create the global cache instance.""" + global _GTP_CACHE + if _GTP_CACHE is None: + _GTP_CACHE = GTPWeightCache() + return _GTP_CACHE + + +def wait_async_comms( + chain_id: str = None, skip_rs: bool = False, finalize_after_drain: bool = False +): + """Drain in-flight GTP async AG / RS handles. + + Inside CUDA graph capture the drains are captured into the graph — the producer-side hook + for cross-graph overlap. A captured cudaStreamWaitEvent on another capture session's event is + a CUDA no-op, so consumers can't wait cross-graph; instead the producer drains here and flags + the param, and the consumer skips its captured wait. + + Args: + chain_id: If specified, only drain params on this chain. + skip_rs: Drain AG only; leave RS in flight. + finalize_after_drain: After RS drain, also accumulate wgrad into + main_grad. Runs main_grad.add_ on rs_stream (right after + NCCL RS) so it starts during AG drain rather than after, + avoiding SM-saturation that blocks cross-graph overlap. + Falls back to caller-stream accumulation if no RS handle. + + Per-param side effects: + * _already_ag_drained = True (if an AG handle was drained) + * _already_finalized = True (if finalize_after_drain=True) + """ + for param in list(_inflight_comm_params): + if ( + chain_id is not None + and getattr(param, "chain_id", GTPChain.UNGRAPHED.value) != chain_id + ): + continue + had_ag = param._prefetch_handle is not None + param._wait_param_gather() + if had_ag: + param._already_ag_drained = True + # Recompute-forward chain: drain its separate in-flight rowwise AG so the + # captured recompute consumer skips its cross-graph wait (full-iteration CG). + if param._recompute_prefetch_handle is not None: + param._wait_recompute_param_gather() + param._recompute_already_drained = True + if not skip_rs: + param._wait_reduce_scatter(finalize_grad=finalize_after_drain) + # Fallback inline-accumulation: only when finalize is requested, _wait_reduce_scatter + # didn't already finalize, and an RS actually ran (rs_ticket set). Skips pure-AG + # prefetches in _inflight_comm_params (no wgrad). + need_fallback_accumulation = ( + finalize_after_drain + and not getattr(param, "_already_finalized", False) + and any(w._rs_ticket is not None for w in param._weights) + ) + if need_fallback_accumulation: + cache = get_global_GTP_cache() + param.rs_event.wait() + for w in param._weights: + w._set_rs_state(GTPWeightState.NONE) + wgrad_rs = cache.get(w._rs_ticket) + w.main_grad.add_(wgrad_rs) + cache.release(w._rs_ticket) + if hasattr(w, "grad_added_to_main_grad"): + w.grad_added_to_main_grad = True + param._already_finalized = True + + +@dataclass +class BatchedNVFP4AllGatherAsyncHandle: + """Handle for batched asynchronous NVFP4 all-gathers.""" + + output_handles: List[_NVFP4AllGatherAsyncHandle] + outer_async_handle: torch.distributed.Work + _synchronized: bool = False + + def wait(self) -> None: + """Wait for the async operation to complete and post-process the tensor.""" + if self._synchronized: + return + self.outer_async_handle.wait() + # Fixes interleaved data for transposed tensor/scale inv and pads scale inv if needed. + for output_handle in self.output_handles: + if output_handle is not None: + assert output_handle.async_handle is None + output_handle.wait() + # release any tensor references just in case + output_handle.output = None + output_handle.columnwise_data_interleaved = None + output_handle.columnwise_scale_inv_interleaved = None + + self._synchronized = True + + +def grouped_gather_along_first_dim( + weights: list, + process_group, + async_op: bool = False, + quantizers: list = None, + output_tensors: list = None, +): + """All-gather multiple weights in one coalesced op; handles NVFP4 post-processing for both + sync and async paths.""" + # Determine device from first weight. + inp = weights[0] + if isinstance(inp, NVFP4TensorStorage): + device = ( + inp._rowwise_data.device + if inp._rowwise_data is not None + else inp._columnwise_data.device + ) + else: + device = inp.device + + weights_all = [] + weight_handles = [] + with torch.distributed._coalescing_manager( + group=process_group, device=device, async_ops=async_op + ) as gather_coalescing_manager: + for i, weight in enumerate(weights): + weight_all, weight_handle = gather_along_first_dim( + weight, + process_group, + quantizer=quantizers[i], + output_tensor=output_tensors[i] if output_tensors is not None else None, + external_coalescing=True, + ) + weights_all.append(weight_all) + weight_handles.append(weight_handle) + + if async_op: + handle = gather_coalescing_manager + has_nvfp4_handles = any(isinstance(wh, _NVFP4AllGatherAsyncHandle) for wh in weight_handles) + if has_nvfp4_handles: + handle = BatchedNVFP4AllGatherAsyncHandle(weight_handles, handle) + else: + for wh in weight_handles: + if isinstance(wh, _NVFP4AllGatherAsyncHandle): + wh.wait() + handle = None + + return weights_all, handle + + +class GTPEmbeddingWeight(torch.autograd.Function): + """All-gather the embedding weight across the GTP group in forward, reduce-scatter its + gradient in backward. + + The weight is stored sharded along the vocab dimension; this materializes the full weight + for the lookup and distributes the gradient back to the shard. + """ + + @staticmethod + def forward(ctx, weight): + """All-gather the full embedding weight across the GTP group for the lookup.""" + ctx.save_for_backward(weight) + return weight.all_gather_and_prefetch(fwd=True) + + @staticmethod + def backward(ctx, grad_output): + """Reduce-scatter the gradient back to this rank's vocab-dim shard.""" + (weight,) = ctx.saved_tensors + return weight.wgrad_reduce_scatter(grad_output) + + +def reset_gtp_state(): + """Clear the process-global GTP prefetch-chain state (GTPShardedParam._chain_state / + ._recompute_chain_state). + + These class-level dicts survive model teardown, so a GTP model rebuilt in-process would + inherit stale last_weight pointers / flushed link tables. Call once before the per-chunk + classify_gtp_chains loop (never inside it — chains span chunks). No-op on a fresh process. + """ + GTPShardedParam._chain_state.clear() + GTPShardedParam._recompute_chain_state.clear() + _GTP_GROUPED_BUF_PARITY_COUNTER.clear() + + +# ------------------------------------------------------------------------ +# Distributed-checkpointing helpers +# ------------------------------------------------------------------------ +# GTP shards axis 0 on top of TP, but the vanilla utils helpers only know TP, so their offsets +# miss the GTP slice. The helper below detects GTPShardedParam per-tensor and composes TP × GTP +# into one axis-0 offset (or two offsets), with replica_id = the DP-with-GTP-with-CP rank. + + +def make_sharded_tensors_for_checkpoint_with_gtp_remat( + state_dict, + prefix, + tensor_parallel_layers_axis_map=None, + sharded_offsets=(), + extra_state_suffix="_extra_state", + *, + tp_group, + dp_cp_group, + intra_dp_cp_group=None, +): + """GTP-aware analogue of make_sharded_tensors_for_checkpoint. + + Per-tensor (is_gtp_param): GTP tensors layer the axis-0 GTP split on the vanilla offsets (FP8 + shards dequantized to BF16 for save); non-GTP tensors delegate to the vanilla helper unchanged, + so this is zero-cost when GTP is inactive. + """ + from megatron.core.transformer.utils import ( # noqa: E402 + make_sharded_object_for_checkpoint, + make_sharded_tensors_for_checkpoint, + ) + from megatron.core.utils import ( # noqa: E402 + get_pg_rank, + get_pg_size, + make_sharded_tensor_for_checkpoint, + make_tp_sharded_tensor_for_checkpoint, + ) + + # Fast path: no GTP-sharded params → defer to vanilla helper, same output. + if not any(is_gtp_param(t) for t in state_dict.values()): + return make_sharded_tensors_for_checkpoint( + state_dict, + prefix, + tensor_parallel_layers_axis_map, + sharded_offsets, + extra_state_suffix=extra_state_suffix, + tp_group=tp_group, + dp_cp_group=dp_cp_group, + ) + + if tensor_parallel_layers_axis_map is None: + tensor_parallel_layers_axis_map = {} + + tp_rank = get_pg_rank(tp_group) + tp_size = get_pg_size(tp_group) + # All GTP params in this state_dict share the same gtp_remat_group (set by the + # wrap hook at module init), so pick it off the first GTP shard. + gtp_remat_group = next(t.group for t in state_dict.values() if is_gtp_param(t)) + gtp_rank = get_pg_rank(gtp_remat_group) + gtp_remat_size = get_pg_size(gtp_remat_group) + + # Replicate-group rank — the true replicas of a given GTP chunk live here. + if intra_dp_cp_group is not None: + dp_replica_rank = get_pg_rank(intra_dp_cp_group) + else: + from megatron.core import parallel_state # noqa: E402 + + dp_replica_rank = parallel_state.get_data_parallel_rank( + with_context_parallel=True, with_gtp_remat=False + ) + + sharded_state_dict = {} + for layer_name, tensor in state_dict.items(): + layer_key = f"{prefix}{layer_name}" + is_gtp_weight_remat = is_gtp_param(tensor) + + if layer_name.endswith(extra_state_suffix): + # ShardedObject (extra_state metadata): GTP-REPLICATED across the GTP group. Fold + # gtp_rank into position 1 of the replica_id (PP, TP-replica-coord, DP) tuple so + # GTP-peer ranks within the same TP slice get unique replica_ids. + replica_id = (0, tp_rank * gtp_remat_size + gtp_rank, dp_replica_rank) + sharded_state_dict[layer_key] = make_sharded_object_for_checkpoint( + tensor, layer_key, sharded_offsets, replica_id=replica_id + ) + continue + + if not is_gtp_weight_remat: + # Non-GTPShardedParam under a GTP-active module (e.g. bias): GTP-replicated, so GTP + # ranks would collide on the same replica_id. Inject gtp_rank into replica_id + # position 1 (same as the GTP-sharded branch below). + if layer_name in tensor_parallel_layers_axis_map: + replica_id = (0, gtp_rank, dp_replica_rank) + sharded_state_dict[layer_key] = make_tp_sharded_tensor_for_checkpoint( + tensor, + layer_key, + tp_axis=tensor_parallel_layers_axis_map[layer_name], + replica_id=replica_id, + prepend_offsets=sharded_offsets, + tp_group=tp_group, + dp_cp_group=dp_cp_group, + ) + else: + replica_id = (0, tp_rank * gtp_remat_size + gtp_rank, dp_replica_rank) + sharded_state_dict[layer_key] = make_sharded_tensor_for_checkpoint( + tensor, + layer_key, + replica_id=replica_id, + prepend_offsets=sharded_offsets, + tp_group=tp_group, + dp_cp_group=dp_cp_group, + ) + continue + + # GTP-sharded tensor: delegate to the GTP-aware single-tensor helper — it layers the + # axis-0 GTP split onto TP, elects the writer over the gtp_remat-excluded DP group, and sets + # allow_shape_mismatch for alignment padding. (tp_axis None → 0; tp_size 1 when no TP.) + tp_axis = tensor_parallel_layers_axis_map.get(layer_name, None) + sharded_state_dict[layer_key] = make_tp_sharded_tensor_for_checkpoint( + tensor, + layer_key, + tp_axis=tp_axis if tp_axis is not None else 0, + prepend_offsets=sharded_offsets, + tp_group=tp_group, + dp_cp_group=dp_cp_group, + ) + + return sharded_state_dict diff --git a/megatron/core/tensor_parallel/gtp_api.py b/megatron/core/tensor_parallel/gtp_api.py new file mode 100644 index 00000000000..b49a5c02ded --- /dev/null +++ b/megatron/core/tensor_parallel/gtp_api.py @@ -0,0 +1,59 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Generalized Tensor Parallelism (GTP) public API. + +Thin re-export of the implementation in +``megatron.core.tensor_parallel.generalized_tensor_parallelism`` (see that module +for the design). GTP depends on TransformerEngine: if TE is missing or too old the +inner module imports cleanly but reports ``HAVE_TE = False``, mirrored here as +``HAVE_GTP = False``. Consumers gate every GTP code path behind ``if HAVE_GTP:``, +so no core module uses GTP symbols without TE. +""" + +try: + from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( + HAVE_TE, + GTPChain, + GTPEmbeddingWeight, + attach_gtp_to_presharded_module, + classify_gtp_remat_chains, + configure_gtp_remat_from_recipe, + dequantize_gtp_native_fp8, + get_ag_stream, + get_rs_stream, + gtp_native_fp8_load_context, + gtp_remat_shard_dim0, + is_gtp_param, + make_sharded_tensors_for_checkpoint_with_gtp_remat, + set_cuda_graph_mempool, + wait_async_comms, + wait_for_gtp_grad_reduction_on_current_stream, + wrap_module_params_gtp, + ) + + HAVE_GTP = HAVE_TE +except ImportError: + # Defensive fallback for any unexpected inner-import failure; consumers import + # the other symbols lazily under an ``if HAVE_GTP:`` guard, so no stubs needed. + HAVE_GTP = False + + +__all__ = [ + "HAVE_GTP", + "GTPChain", + "GTPEmbeddingWeight", + "attach_gtp_to_presharded_module", + "classify_gtp_remat_chains", + "configure_gtp_remat_from_recipe", + "dequantize_gtp_native_fp8", + "get_ag_stream", + "get_rs_stream", + "gtp_native_fp8_load_context", + "gtp_remat_shard_dim0", + "is_gtp_param", + "make_sharded_tensors_for_checkpoint_with_gtp_remat", + "set_cuda_graph_mempool", + "wait_async_comms", + "wait_for_gtp_grad_reduction_on_current_stream", + "wrap_module_params_gtp", +] diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index c072c52bd05..3ff635b362b 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -16,10 +16,13 @@ from megatron.core.model_parallel_config import ModelParallelConfig from megatron.core.parallel_state import ( + get_expert_gtp_weight_remat_rank, get_global_memory_buffer, + get_gtp_weight_remat_rank, get_tensor_model_parallel_rank, get_tensor_model_parallel_world_size, ) +from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.utils import ( divide, get_pg_rank, @@ -89,11 +92,17 @@ dist_reduce_scatter_func = torch.distributed._reduce_scatter_base -def param_is_not_tensor_parallel_duplicate(param, tp_group=None): - """Returns true if the passed-in parameter is not a duplicate parameter - on another TP rank.""" +def param_is_not_tensor_parallel_duplicate(param, tp_group=None, expert_tp_group=None): + """Return whether a parameter contributes to a unique model-parallel shard. + + Parameters reduced over expert data parallel groups use the expert tensor-parallel + group for duplicate filtering. Other parameters use the regular tensor-parallel group. + """ if hasattr(param, "tensor_model_parallel") and param.tensor_model_parallel: return True + # allreduce=False marks parameters reduced over expert DP, so filter their duplicates over ETP. + if not getattr(param, "allreduce", True) and expert_tp_group is not None: + tp_group = expert_tp_group # Prefer provided tp_group when available (new explicit path). if tp_group is not None: return tp_group.rank() == 0 @@ -101,6 +110,29 @@ def param_is_not_tensor_parallel_duplicate(param, tp_group=None): return get_tensor_model_parallel_rank() == 0 +def copy_gtp_attributes(destination, source): + """Copy the GTP dedup tags (is_gtp_weight_remat, allreduce) onto a param view/copy, so the + optimizer's master shards stay classifiable by param_is_not_gtp_duplicate.""" + for attr in ("is_gtp_weight_remat", "allreduce"): + if hasattr(source, attr): + setattr(destination, attr, getattr(source, attr)) + + +def param_is_not_gtp_duplicate(param): + """True if the param's grad is counted once across the GTP_remat/EGTP_remat axis. + + GTP_remat/EGTP_remat shards are unique per peer (kept); replicated params counted only on + rank 0 of the gtp_remat/egtp_remat axis (else counted N times). When GTP_remat is off rank is 0, + so every param is kept. + """ + if getattr(param, "is_gtp_weight_remat", False): + return True + is_expert = not getattr(param, "allreduce", True) + if is_expert: + return get_expert_gtp_weight_remat_rank() == 0 + return get_gtp_weight_remat_rank() == 0 + + def set_tensor_model_parallel_attributes(tensor, is_parallel, dim, stride): """Sets tp attributes to tensor""" # Make sure the attributes are not set. @@ -281,6 +313,20 @@ def __init__( tensor=self.weight, is_parallel=True, dim=0, stride=1 ) + self.gtp_remat_size = 1 + gtp_remat_group = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat"] + ).gtp_remat + if gtp_remat_group is not None and gtp_remat_group.size() > 1: + from megatron.core.tensor_parallel.gtp_api import wrap_module_params_gtp + + wrap_module_params_gtp(self, ["weight"], gtp_remat_group) + self.gtp_remat_size = gtp_remat_group.size() + # Nothing prefetches embedding — it is head of the UNGRAPHED + # chain in fwd, and its bwd bypasses all_gather_and_prefetch_bwd + # via GTPEmbeddingWeight.backward. + self.weight._need_weight_prefetch = False + def forward(self, input_): """Forward. @@ -295,12 +341,19 @@ def forward(self, input_): masked_input[input_mask] = 0 else: masked_input = input_ + + weight = self.weight + if self.gtp_remat_size > 1: + from megatron.core.tensor_parallel.gtp_api import GTPEmbeddingWeight + + weight = GTPEmbeddingWeight.apply(self.weight) + # Get the embeddings. if self.deterministic_mode: - output_parallel = self.weight[masked_input] + output_parallel = weight[masked_input] else: # F.embedding currently has a non-deterministic backward function - output_parallel = F.embedding(masked_input, self.weight) + output_parallel = F.embedding(masked_input, weight) # Mask the output embedding. if self.tp_group.size() > 1: output_parallel[input_mask, :] = 0.0 @@ -400,6 +453,7 @@ def linear_with_frozen_weight( tp_group: Optional[torch.distributed.ProcessGroup], grad_output_buffer: Optional[List[torch.Tensor]] = None, wgrad_deferral_limit: None = None, + gtp_remat_size: int = 1, ) -> torch.Tensor: """Linear layer execution with weight.requires_grad == False. @@ -436,6 +490,10 @@ def linear_with_frozen_weight( wgrad_deferral_limit (int optional): dummy argument, used to keep the API unified between all forward implementation functions. + + gtp_remat_size (int): GTP shard count. When > 1 the weight is GTP-sharded and must be + all-gathered to its full shape before the matmul, mirroring the trainable path. + Defaults to 1 (no-op) for the common non-GTP / non-sharded case. """ assert grad_output_buffer is None, ( @@ -456,6 +514,9 @@ def linear_with_frozen_weight( else: input = input + if gtp_remat_size > 1: + weight = weight.all_gather_and_prefetch(fwd=True) + args = [input, weight, bias, allreduce_dgrad, tp_group] return LinearWithFrozenWeight.apply(*args) @@ -477,6 +538,7 @@ def forward( grad_output_buffer, wgrad_deferral_limit, tp_group, + gtp_remat_size, ): """Forward.""" if gradient_accumulation_fusion and hasattr(weight, "main_grad"): @@ -484,6 +546,10 @@ def forward( else: main_grad = None ctx.save_for_backward(input, weight) + + if gtp_remat_size > 1: + weight = weight.all_gather_and_prefetch(fwd=True) + # We can't save main_grad in save_for_backward as this module would be # reused across layers like MTP logits. So, to prevent in-place modification # checks we save the tensor in ctx. @@ -495,6 +561,7 @@ def forward( ctx.wgrad_deferral_limit = wgrad_deferral_limit ctx.grad_output_buffer = grad_output_buffer ctx.tp_group = tp_group + ctx.gtp_remat_size = gtp_remat_size if sequence_parallel: dim_size = list(input.size()) @@ -518,6 +585,13 @@ def backward(ctx, grad_output): input, weight = ctx.saved_tensors main_grad = ctx.main_grad use_bias = ctx.use_bias + + # GTP: re-gather weight for dgrad + if ctx.gtp_remat_size > 1: + sharded_weight = weight + weight = sharded_weight.all_gather_and_prefetch_bwd() + ctx.gradient_accumulation_fusion = False + grad_output_buffer = ctx.grad_output_buffer wgrad_deferral_limit = ctx.wgrad_deferral_limit handle = None @@ -651,16 +725,31 @@ def backward(ctx, grad_output): grad_weight = grad_output.t().matmul(total_input) grad_bias = grad_output.sum(dim=0) if use_bias else None + # GTP: reduce-scatter wgrad + if ctx.gtp_remat_size > 1 and grad_weight is not None: + grad_weight = sharded_weight.wgrad_reduce_scatter(grad_weight) + if ctx.sequence_parallel: handle.wait() # Need to return None's as gradient has to flow for all the input arguments # provided during forward - return (sub_grad_input, grad_weight, grad_bias, None, None, None, None, None, None) + return ( + sub_grad_input, + grad_weight, + grad_bias, + None, + None, + None, + None, + None, + None, + None, + ) if ctx.allreduce_dgrad: handle.wait() - return grad_input, grad_weight, grad_bias, None, None, None, None, None, None + return grad_input, grad_weight, grad_bias, None, None, None, None, None, None, None def linear_with_grad_accumulation_and_async_allreduce( @@ -673,6 +762,7 @@ def linear_with_grad_accumulation_and_async_allreduce( grad_output_buffer: Optional[List[torch.Tensor]] = None, wgrad_deferral_limit: Optional[int] = 0, tp_group: Optional[torch.distributed.ProcessGroup] = None, + gtp_remat_size: int = 1, ) -> torch.Tensor: """Linear layer execution with asynchronous communication and gradient accumulation fusion in backprop. @@ -749,6 +839,7 @@ def linear_with_grad_accumulation_and_async_allreduce( grad_output_buffer, wgrad_deferral_limit, tp_group, + gtp_remat_size, ] if not linear_with_grad_accumulation_and_async_allreduce.warned: @@ -867,6 +958,10 @@ def __init__( world_size = get_pg_size(self.tp_group) rank = get_pg_rank(self.tp_group) self.explicit_expert_comm = self.is_expert and (world_size > 1 or self.expert_parallel) + use_expert_pgs = self.is_expert and ( + self.expert_parallel + or self.config.expert_tensor_parallel_size != self.config.tensor_model_parallel_size + ) self.output_size_per_partition = divide(output_size, world_size) # Parameters. @@ -919,10 +1014,21 @@ def __init__( tensor=self.weight, is_parallel=True, dim=0, stride=stride ) - setattr(self.weight, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr(self.weight, "allreduce", not use_expert_pgs) else: self.weight = None + self.gtp_remat_size = 1 + _pg = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat", "expt_gtp_remat"] + ) + gtp_remat_group = _pg.expt_gtp_remat if self.is_expert else _pg.gtp_remat + if gtp_remat_group is not None and gtp_remat_group.size() > 1: + from megatron.core.tensor_parallel.gtp_api import wrap_module_params_gtp + + wrap_module_params_gtp(self, ["weight"], gtp_remat_group) + self.gtp_remat_size = gtp_remat_group.size() + if bias: if config.use_cpu_initialization: self.bias = Parameter( @@ -941,7 +1047,7 @@ def __init__( # Always initialize bias to zero. with torch.no_grad(): self.bias.zero_() - setattr(self.bias, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr(self.bias, "allreduce", not use_expert_pgs) else: self.register_parameter("bias", None) @@ -1075,6 +1181,7 @@ def forward( else None ), tp_group=self.tp_group, + gtp_remat_size=self.gtp_remat_size, ) gather_output = self.gather_output @@ -1269,7 +1376,22 @@ def __init__( set_tensor_model_parallel_attributes( tensor=self.weight, is_parallel=True, dim=1, stride=stride ) - setattr(self.weight, "allreduce", not (self.is_expert and self.expert_parallel)) + use_expert_pgs = self.is_expert and ( + self.expert_parallel + or self.config.expert_tensor_parallel_size != self.config.tensor_model_parallel_size + ) + setattr(self.weight, "allreduce", not use_expert_pgs) + + self.gtp_remat_size = 1 + _pg = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat", "expt_gtp_remat"] + ) + gtp_remat_group = _pg.expt_gtp_remat if self.is_expert else _pg.gtp_remat + if gtp_remat_group is not None and gtp_remat_group.size() > 1: + from megatron.core.tensor_parallel.gtp_api import wrap_module_params_gtp + + wrap_module_params_gtp(self, ["weight"], gtp_remat_group) + self.gtp_remat_size = gtp_remat_group.size() if bias: if config.use_cpu_initialization: @@ -1287,7 +1409,7 @@ def __init__( # Always initialize bias to zero. with torch.no_grad(): self.bias.zero_() - setattr(self.bias, "allreduce", not (self.is_expert and self.expert_parallel)) + setattr(self.bias, "allreduce", not use_expert_pgs) setattr(self.bias, "sequence_parallel", self.sequence_parallel) else: self.register_parameter("bias", None) @@ -1343,6 +1465,7 @@ def forward(self, input_: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: sequence_parallel=False, tp_group=None, grad_output_buffer=None, + gtp_remat_size=self.gtp_remat_size, ) # All-reduce across all the partitions. diff --git a/megatron/core/tensor_parallel/random.py b/megatron/core/tensor_parallel/random.py index dbecd73abd3..7da75366169 100644 --- a/megatron/core/tensor_parallel/random.py +++ b/megatron/core/tensor_parallel/random.py @@ -18,8 +18,12 @@ from typing_extensions import TypeVarTuple, Unpack from megatron.core.parallel_state import ( + get_expert_gtp_weight_remat_rank, + get_expert_gtp_weight_remat_world_size, get_expert_model_parallel_rank, get_expert_tensor_parallel_rank, + get_gtp_weight_remat_rank, + get_gtp_weight_remat_world_size, get_tensor_model_parallel_rank, ) from megatron.core.utils import is_te_min_version, safely_set_viewless_tensor_data @@ -91,6 +95,10 @@ def _get_share_storage(): _MODEL_PARALLEL_RNG_TRACKER_NAME = 'model-parallel-rng' _EXPERT_PARALLEL_RNG_TRACKER_NAME = 'expert-parallel-rng' _DATA_PARALLEL_RNG_TRACKER_NAME = 'data-parallel-rng' +# GTP_remat weight-init trackers: shards init per-rank, so each peer must draw DIFFERENT values; +# registered only when the axis is active (see model_parallel_cuda_manual_seed). +_GTP_REMAT_RNG_TRACKER_NAME = 'gtp-remat-rng' +_EXPERT_GTP_REMAT_RNG_TRACKER_NAME = 'egtp-remat-rng' def _get_cuda_rng_state( @@ -213,6 +221,11 @@ def get_data_parallel_rng_tracker_name(): return _DATA_PARALLEL_RNG_TRACKER_NAME +def get_gtp_remat_rng_tracker_name(is_expert=False): + """Get the (E)GTP_remat weight-init rng tracker name (per-(E)GTP-rank distinct draws).""" + return _EXPERT_GTP_REMAT_RNG_TRACKER_NAME if is_expert else _GTP_REMAT_RNG_TRACKER_NAME + + class CudaRNGStatesTracker: """Tracker for the cuda RNG states. @@ -483,6 +496,19 @@ def model_parallel_cuda_manual_seed( expert_parallel_seed = seed + 1024 + 100 * ep_rank + etp_rank _CUDA_RNG_STATE_TRACKER.add(_EXPERT_PARALLEL_RNG_TRACKER_NAME, expert_parallel_seed) + # GTP_remat weight-init states: shards are initialized per-rank (GTP-agnostic init), so peers + # must draw DIFFERENT values (everything above is identical across peers by design). The 65536 + # stride keeps these disjoint from the tp/ep/etp seeds. Added only when the axis is active, so + # non-GTP runs keep a byte-identical tracker set (and checkpoint rng payload). + gtp_remat_rank = get_gtp_weight_remat_rank() + if get_gtp_weight_remat_world_size() > 1: + gtp_remat_seed = tensor_model_parallel_seed + 65536 * (1 + gtp_remat_rank) + _CUDA_RNG_STATE_TRACKER.add(_GTP_REMAT_RNG_TRACKER_NAME, gtp_remat_seed) + egtp_remat_rank = get_expert_gtp_weight_remat_rank() + if get_expert_gtp_weight_remat_world_size() > 1: + egtp_remat_seed = expert_parallel_seed + 32768 + 65536 * (1 + egtp_remat_rank) + _CUDA_RNG_STATE_TRACKER.add(_EXPERT_GTP_REMAT_RNG_TRACKER_NAME, egtp_remat_seed) + def is_graph_safe_cuda_rng_tracker(cuda_rng_tracker): """Check if the cuda rng tracker is graph safe version.""" diff --git a/megatron/core/tokenizers/text/parsers/__init__.py b/megatron/core/tokenizers/text/parsers/__init__.py index dc27763f905..d541cb4a74d 100644 --- a/megatron/core/tokenizers/text/parsers/__init__.py +++ b/megatron/core/tokenizers/text/parsers/__init__.py @@ -2,11 +2,15 @@ from megatron.core.tokenizers.text.parsers.deepseek_r1_reasoning_parser import ( DeepSeekR1ReasoningParser, ) +from megatron.core.tokenizers.text.parsers.nemotron_v3_reasoning_parser import ( + NemotronV3ReasoningParser, +) from megatron.core.tokenizers.text.parsers.qwen3_coder_tool_parser import Qwen3CoderToolParser PARSER_MAPPING = { "deepseek-r1-reasoning": DeepSeekR1ReasoningParser, "qwen3-coder-tool": Qwen3CoderToolParser, + "nemotron-v3-reasoning": NemotronV3ReasoningParser, } __all__ = ["PARSER_MAPPING"] diff --git a/megatron/core/tokenizers/text/parsers/nemotron_v3_reasoning_parser.py b/megatron/core/tokenizers/text/parsers/nemotron_v3_reasoning_parser.py new file mode 100644 index 00000000000..c5468059240 --- /dev/null +++ b/megatron/core/tokenizers/text/parsers/nemotron_v3_reasoning_parser.py @@ -0,0 +1,62 @@ +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +from megatron.core.tokenizers.text.parsers.deepseek_r1_reasoning_parser import ( + DeepSeekR1ReasoningParser, +) + + +class NemotronV3ReasoningParser(DeepSeekR1ReasoningParser): + """Parser for NVIDIA Nemotron 3 (Super, Ultra) reasoning output. + + Behaves like `DeepSeekR1ReasoningParser`, except when reasoning is disabled + via `enable_thinking=False`, or the caller passes `force_nonempty_content=True`: + in that case, if no content would otherwise be returned (either because + `` never closes, e.g. reasoning exceeded the max length, or because + it closes with nothing following it), the reasoning text is returned as + content instead of being discarded, so callers always get a non-empty + response. + """ + + @staticmethod + def _should_force_content(chat_template_kwargs: "dict | None") -> bool: + """Whether would-be-empty content should be backfilled from reasoning. + + Mirrors vLLM's `SuperV3ReasoningParser._should_force_content`: force + content when reasoning is disabled (`enable_thinking is False`) or the + caller explicitly requests it (`force_nonempty_content is True`). Both + flags are supplied by the client inside `chat_template_kwargs`. + """ + return bool( + chat_template_kwargs + and ( + chat_template_kwargs.get("enable_thinking") is False + or chat_template_kwargs.get("force_nonempty_content") is True + ) + ) + + @staticmethod + def parse(text: str, **kwargs) -> tuple[str, dict[str, str]]: + """Extract reasoning content delimited by `...` tags. + + Delegates the ``/`` split to `DeepSeekR1ReasoningParser`, + then surfaces the reasoning as content (instead of discarding it) when + reasoning was disabled or the caller forced non-empty content. + + Args: + text (str): The text to parse. + chat_template_kwargs (dict, optional): The request's + `chat_template_kwargs`. When it sets `enable_thinking=False` or + `force_nonempty_content=True`, reasoning is surfaced as content + rather than discarded if there would otherwise be no content. + + Returns: + tuple[str, dict[str, str]]: A tuple containing the unprocessed text + and a dictionary with the extracted reasoning content. + """ + content, info = DeepSeekR1ReasoningParser.parse(text, **kwargs) + if ( + content == "" + and info.get("reasoning") + and NemotronV3ReasoningParser._should_force_content(kwargs.get("chat_template_kwargs")) + ): + return info["reasoning"], {} + return content, info diff --git a/megatron/core/transformer/attention.py b/megatron/core/transformer/attention.py index 1f29c93eef3..96f4b9d6fa6 100644 --- a/megatron/core/transformer/attention.py +++ b/megatron/core/transformer/attention.py @@ -88,10 +88,22 @@ except ImportError as e: pass +# The FA4 version is tracked by the `flash-attn-4` distribution metadata, +# not `flash_attn.__version__` (which reports the 2.x version) or +# `flash_attn.cute.__version__` (which is 0.0.0), so we cannot use +# `is_fa_min_version` here. +_MIN_FA4_VERSION = "4.0.0b20" try: + from importlib.metadata import PackageNotFoundError + from importlib.metadata import version as _get_dist_version + from flash_attn.cute import flash_attn_varlen_func as flash_attn4_varlen_func + from packaging.version import Version as _Version - HAVE_FA4 = True + try: + HAVE_FA4 = _Version(_get_dist_version("flash-attn-4")) >= _Version(_MIN_FA4_VERSION) + except PackageNotFoundError: + HAVE_FA4 = False except ImportError: HAVE_FA4 = False @@ -314,6 +326,7 @@ def __init__( self.attn_mask_type = attn_mask_type self.attention_type = attention_type self.batch_invariant_mode = config.batch_invariant_mode + self.flash_attention_version = config.flash_attention_version # Cache the YaRN concentration factor (a.k.a. attention factor / mscale), # which is a pure function of the config and is reused on every forward @@ -1061,6 +1074,29 @@ def _flash_attention_3_forward_wrapper( ) return output_total, softmax_lse + def _resolve_flash_version(self) -> Tuple[bool, bool]: + """Resolve which FlashAttention generation this attention should run. + + Honors ``config.flash_attention_version`` when pinned, otherwise falls back + to the auto preference order (FA4 > FA3 > FA2). Returns ``(use_fa4, use_fa3)``; + when both are False the FA2 kernel is used. + """ + pinned = self.flash_attention_version + if pinned == 4: + assert ( + HAVE_FA4 + ), "flash_attention_version=4 requested but FlashAttention-4 is not installed" + return True, False + if pinned == 3: + assert ( + HAVE_FA3 + ), "flash_attention_version=3 requested but FlashAttention-3 is not installed" + return False, True + if pinned == 2: + return False, False + # Auto: prefer the newest available generation. + return HAVE_FA4, (HAVE_FA3 and not HAVE_FA4) + def flash_decode_and_prefill( self, q: Tensor, @@ -1119,6 +1155,8 @@ def flash_decode_and_prefill( # the sink (off-by-one / learnable) softmax correction post-hoc. need_lse = softmax_offset is not None + use_fa4, use_fa3 = self._resolve_flash_version() + # Flash attn kernel. if not is_decode_only: q = q.squeeze(1) @@ -1126,7 +1164,7 @@ def flash_decode_and_prefill( softmax_scale = self.softmax_scale else: softmax_scale = q.shape[-1] ** -0.5 - if HAVE_FA4: + if use_fa4: output_total, softmax_lse = flash_attn4_varlen_func( q, k, @@ -1139,9 +1177,9 @@ def flash_decode_and_prefill( softmax_scale=softmax_scale, causal=True, window_size=window_size, - num_splits=1, + num_splits=0 if not self.batch_invariant_mode else 1, ) - elif HAVE_FA3: + elif use_fa3: # TODO(ksanthanam): Replace with call to flash_attn_varlen_func once # it accepts block_table fa3_ret = self._flash_attention_3_forward_wrapper( @@ -1242,7 +1280,7 @@ def flash_decode_and_prefill( output_total, softmax_lse, softmax_offset ) else: - if HAVE_FA4: + if use_fa4: if getattr(self, "softmax_scale", None) is not None: softmax_scale = self.softmax_scale else: @@ -1261,7 +1299,7 @@ def flash_decode_and_prefill( softmax_scale=softmax_scale, causal=True, window_size=window_size, - num_splits=1, + num_splits=0 if not self.batch_invariant_mode else 1, ) if need_lse: # output_total: (B*S, H, D); softmax_lse: (H, B*S) @@ -1285,12 +1323,12 @@ def flash_decode_and_prefill( "softmax_scale": softmax_scale, "causal": True, "window_size": window_size, - "page_table" if HAVE_FA3 else "block_table": block_table, + "page_table" if use_fa3 else "block_table": block_table, "num_splits": 0 if not self.batch_invariant_mode else 1, } if need_lse: flash_attn_args["return_softmax_lse"] = True - if HAVE_FA3: + if use_fa3: kvcache_ret = flash_attn3_with_kvcache(**flash_attn_args) else: assert ( diff --git a/megatron/core/transformer/cuda_graphs.py b/megatron/core/transformer/cuda_graphs.py index 00ec8f6ebc4..139b0b1bb3e 100644 --- a/megatron/core/transformer/cuda_graphs.py +++ b/megatron/core/transformer/cuda_graphs.py @@ -59,6 +59,31 @@ except: HAVE_TE_GRAPHS = False +try: + from megatron.core.tensor_parallel.gtp_api import HAVE_GTP +except ImportError: + # GTP requires TransformerEngine with the GTP hook registry; treat it as + # unavailable when that import path cannot be resolved. + HAVE_GTP = False + +if HAVE_GTP: + from megatron.core.tensor_parallel.gtp_api import ( + GTPChain, + get_ag_stream, + get_rs_stream, + set_cuda_graph_mempool, + wait_async_comms, + ) +else: + # Placeholders so static analysis does not flag these GTP-only symbols as + # possibly-used-before-assignment; every use site is guarded by HAVE_GTP / + # gtp_remat at runtime. + GTPChain = None + get_ag_stream = None + get_rs_stream = None + set_cuda_graph_mempool = None + wait_async_comms = None + try: from tqdm import tqdm @@ -71,6 +96,59 @@ logger = logging.getLogger(__name__) +def _get_tensor_alias_chain(tensor): + """Return a tensor followed by each underlying base tensor.""" + aliases = [] + while torch.is_tensor(tensor): + aliases.append(tensor) + base = getattr(tensor, "_base", None) + if base is None or base is tensor: + break + tensor = base + return aliases + + +def _apply_cudagraph_buffer_metadata(tensor, *, is_output=False): + """Attach one shared CUDA graph metadata object to a tensor and its base chain.""" + aliases = _get_tensor_alias_chain(tensor) + metadata = next( + (alias.cg_buffer_metadata for alias in aliases if hasattr(alias, "cg_buffer_metadata")), + None, + ) + if is_output: + metadata = CudagraphBufferMetadata( + is_cudagraph_output=True, + is_saved_for_backward=bool(metadata and metadata.is_saved_for_backward), + ) + elif metadata is None: + metadata = CudagraphBufferMetadata() + for alias in aliases: + alias.cg_buffer_metadata = metadata + return metadata + + +def _tag_cudagraph_buffer_saved_for_backward(tensor): + """Tag a CUDA graph input or output observed in a Python 'save_for_backward' call.""" + if not torch.is_tensor(tensor): + return + + # Views of the same graph buffer share one metadata object. If this tensor has not reached a + # graph boundary yet, initialize its metadata now so record-time input/output classification + # can preserve the saved-for-backward lifetime. + metadata = _apply_cudagraph_buffer_metadata(tensor) + metadata.is_saved_for_backward = True + + +_GTP_RUNNER_STREAMS: List[torch.cuda.Stream] = [] + + +def get_gtp_runner_streams() -> List[torch.cuda.Stream]: + """Replay streams of all GTP CG runners; finalize_model_grads waits on these + (tail = captured Phase 2 main_grad.add_) before reading main_grad. + """ + return _GTP_RUNNER_STREAMS + + def _set_skip_fp8_weight_update_tensor(skip: bool) -> None: """Toggle TE's FP8 "skip weight refresh" flag between microbatches. @@ -148,10 +226,25 @@ class CudagraphBufferMetadata: Metadata saved to tensors during cudagraph capture. This data will be used to determine during graph captue when a cudagraph can reuse a buffer or directly write its output into a subsequent's graph's input. + + Set during recording: + is_cudagraph_input / is_cudagraph_output — which graph boundary this buffer sits on. + is_saved_for_backward — set by the save_for_backward observer; means the forward + buffer must outlive the forward graph and stay allocator-owned until backward + capture. + + Reuse accounting (used during graph creation): + input_use_count — times this buffer appears as a graph input. + cudagraph_reuse_ref_count / capture_reuse_count — remaining reuses; drives + can_skip_replay_copy and when args_to_clear_buffers fires. + fwd_cudagraph_buffer / bwd_cudagraph_buffer — the shared strong-ref buffer other + graphs alias for this input/grad. + """ is_cudagraph_input: bool = False is_cudagraph_output: bool = False + is_saved_for_backward: bool = False input_use_count: int = 0 cudagraph_reuse_ref_count: int = 0 capture_reuse_count: int = 0 @@ -210,7 +303,19 @@ def wrapper(arg): changes = { f.name: tree_map_pyt(func, getattr(arg, f.name)) for f in dataclasses.fields(arg) } - return dataclasses.replace(arg, **changes) + mapped_arg = dataclasses.replace(arg, **changes) + + # 'dataclasses.replace' reruns '__post_init__', which may overwrite a tensor + # field that was explicitly mapped above. In particular, PackedSeqParams rebuilds + # 'seq_idx' from 'cu_seqlens'. CUDA graph input buffers are zero-initialized, so + # that rebuild assigns every token the padded sequence count and can make Mamba + # kernels access out of bounds during graph capture. Preserve the tensor selected by + # the mapping operation; replay will populate that buffer with the real input value. + for name, value in changes.items(): + if torch.is_tensor(value) and getattr(mapped_arg, name) is not value: + object.__setattr__(mapped_arg, name, value) + + return mapped_arg # Otherwise, apply the user function return func(arg) @@ -342,6 +447,36 @@ def create_strong_ref(ten: torch.Tensor): bwd_buffer_reuse_ref_count = 0 +def _backup_grads_before_capture(runner): + """Snapshot main_grad so create_fwd_graph's eager warmup can't corrupt the finalized grads; + restore with '_restore_grads_after_capture'. + """ + backup = {} + for p in runner.base_module.parameters(): + mg = getattr(p, "main_grad", None) + if mg is not None: + backup[id(p)] = (p, mg.clone()) + + if runner.gtp_remat: + # GTP only: also protect the cross-graph next_w the cascade accumulates into. + for p in runner.base_module.parameters(): + nw = getattr(p, "next_w", None) if getattr(p, "is_gtp_weight_remat", False) else None + if nw is None: + continue + shards = nw.weight_list if getattr(nw, "is_routed_expert", False) else [nw] + for w in shards or []: + mg = getattr(w, "main_grad", None) + if mg is not None and id(w) not in backup: + backup[id(w)] = (w, mg.clone()) + return backup + + +def _restore_grads_after_capture(backup): + """Restore the main_grad snapshots taken by '_backup_grads_before_capture'.""" + for p, saved in backup.values(): + p.main_grad.copy_(saved) + + class _CudagraphGlobalRecord: """A global datastructure that records of the ordering of all _CudaGraphRunner's first fwd or bwd passes. 'create_cudagraphs' will use this to create @@ -355,6 +490,36 @@ class _CudagraphGlobalRecord: 'record_bwd_graph.""" cudagraph_record: list[tuple] = [] cudagraph_inference_record: list[tuple] = [] + _saved_tensors_observer = None + + @classmethod + def _enable_saved_tensors_observer(cls): + """Observe Python 'save_for_backward' calls while recording and capturing graphs.""" + if cls.cudagraph_created or cls._saved_tensors_observer is not None: + return + + function_ctx = torch.autograd.function.FunctionCtx + original_save_for_backward = function_ctx.save_for_backward + + def observing_save_for_backward(ctx, *tensors): + for tensor in tensors: + _tag_cudagraph_buffer_saved_for_backward(tensor) + return original_save_for_backward(ctx, *tensors) + + cls._saved_tensors_observer = (original_save_for_backward, observing_save_for_backward) + function_ctx.save_for_backward = observing_save_for_backward + + @classmethod + def _disable_saved_tensors_observer(cls): + """Restore Python's original 'save_for_backward' implementation.""" + if cls._saved_tensors_observer is None: + return + + original_save_for_backward, observing_save_for_backward = cls._saved_tensors_observer + function_ctx = torch.autograd.function.FunctionCtx + if function_ctx.save_for_backward is observing_save_for_backward: + function_ctx.save_for_backward = original_save_for_backward + cls._saved_tensors_observer = None @classmethod def record_fwd_graph(cls, runner, args, kwargs, out): @@ -368,6 +533,14 @@ def record_bwd_graph(cls, runner): @classmethod def create_cudagraphs(cls): + """Create recorded CUDA graphs, then remove the saved-tensor observer.""" + try: + return cls._create_cudagraphs() + finally: + cls._disable_saved_tensors_observer() + + @classmethod + def _create_cudagraphs(cls): """Iterate through 'cudagraph_record' creating graphs in the order in which they were recorded.""" # Cudagraphs have already been created, check that no cudagraphed modules ran in eager mode @@ -510,6 +683,8 @@ def create_cudagraphs(): def delete_cuda_graphs(): """Delete all CUDA graphs.""" + _CudagraphGlobalRecord._disable_saved_tensors_observer() + # Reset runners. for record in [ *_CudagraphGlobalRecord.cudagraph_record, @@ -529,6 +704,7 @@ def delete_cuda_graphs(): _CudagraphGlobalRecord.cudagraph_created = False _CudagraphGlobalRecord.cudagraph_record = [] _CudagraphGlobalRecord.cudagraph_inference_record = [] + _GTP_RUNNER_STREAMS.clear() # TODO: Optional?: Force garbage collection to clean up memory gc.collect() @@ -599,7 +775,12 @@ def forward(ctx, runner, is_first_microbatch, *inputs): can_skip_replay_copy = getattr( cudagraph_input, "can_skip_replay_copy", False ) and getattr(user_input, "can_skip_replay_copy", True) - if can_skip_replay_copy: + + # When the same input (like cu_seqlens) is passed to multiple cudagraphs, the first + # cudagraph copies it into the corresponding 'cudagraph_input'. Subsequent cudagraphs + # will then read the same cudagraph_input, leading to a case where the passed tensor + # doesn't need a copy despite being a different data_ptr as its 'cudagraph_input'. + if can_skip_replay_copy and cudagraph_input.cg_buffer_metadata.input_use_count == 1: assert user_input.data_ptr() == cudagraph_input.data_ptr() elif user_input.data_ptr() != cudagraph_input.data_ptr(): cudagraph_input.copy_(user_input) @@ -625,13 +806,13 @@ def forward(ctx, runner, is_first_microbatch, *inputs): _set_skip_fp8_weight_update_tensor(not is_first_microbatch) runner.fp8_param_cache_updated = is_first_microbatch - runner.fwd_graph.replay() - - if runner.is_last_layer: - outputs = tuple(torch.clone(t) for t in runner.fwd_graph_output_surface) - for output in outputs: - output.can_skip_replay_copy = False - return outputs + if runner.use_stream: + runner.stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(runner.stream): + runner.fwd_graph.replay() + torch.cuda.current_stream().wait_event(runner.fwd_completion_event) + else: + runner.fwd_graph.replay() return runner.fwd_graph_output_surface @staticmethod @@ -656,7 +837,14 @@ def backward(ctx, *grads): if user_output_grad.data_ptr() != cudagraph_output_grad.data_ptr(): cudagraph_output_grad.copy_(user_output_grad) - runner.bwd_graph.replay() + if runner.use_stream: + runner.stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(runner.stream): + runner.bwd_graph.replay() + torch.cuda.current_stream().wait_event(runner.bwd_completion_event) + else: + runner.bwd_graph.replay() + runner.bwd_graph_replay_complete_event.record(torch.cuda.current_stream()) for param in runner.params_to_backprop: param._cudagraph_wgrad_ready_event = runner.bwd_graph_replay_complete_event @@ -668,6 +856,20 @@ def backward(ctx, *grads): ): FP8GlobalStateManager.reduce_and_update_fp8_tensors(forward=False) + # DDP grad-ready hook is silenced at capture/replay, so fire it here (on each param's + # rs_stream, after wait_stream(runner.stream) fences Phase 2) to let DDP RS overlap bwd. + if runner.gtp_remat: + for gtp_rs_stream, params in runner._gtp_finalize_hook_plan: + gtp_rs_stream.wait_stream(runner.stream) + with torch.cuda.stream(gtp_rs_stream): + for param in params: + param.grad = None + if hasattr(param, 'grad_added_to_main_grad'): + param.grad_added_to_main_grad = True + hook = getattr(param, '_grad_accum_hook', None) + if hook is not None: + hook() + return None, None, *runner.static_grad_inputs, *(None,) * len(runner.params_to_backprop) @@ -715,6 +917,16 @@ def __init__( self.fp4_runtime_enabled = None self.deallocate_pipeline_outputs = False self.num_warmup_steps = 0 + self.use_stream = False + self.gtp_remat = False + self.fwd_side_streams = [] + self.bwd_side_streams = [] + # Populated by create_bwd_graph: GTP params whose main_grad.add_ was captured in THIS + # graph. Used in Graphed.backward's post-replay hook loop to fire DDP hooks only in the + # graph whose replay populates main_grad. + self.finalized_during_bwd_capture = [] + # (rs_stream, params) DDP grad-ready hook plan; built in create_bwd_graph. + self._gtp_finalize_hook_plan = [] self.grad_enabled = need_backward and torch.is_grad_enabled() self.func = super(MegatronModule, self.base_module).__call__ if func is None else func @@ -737,6 +949,31 @@ def __init__( self.fp4_enabled = self.base_module.config.fp4 is not None self.fp8_runtime_enabled = None self.fp4_runtime_enabled = None + self.gtp_remat = self.base_module.config.gtp_weight_remat_size > 1 + + if self.gtp_remat: + # Ensure internal warmup (inside create_fwd_graph) has >= 2 steps + # for GTP: 1st builds chain + tickets, 2nd exercises prefetch path. + self.num_warmup_steps = max(self.num_warmup_steps, 2) + + self.use_stream = True + self.stream = torch.cuda.Stream() + self.fwd_completion_event = torch.cuda.Event(external=True, interprocess=True) + self.bwd_completion_event = torch.cuda.Event(external=True, interprocess=True) + # Register (chain, group) side streams before the first forward. + # Dense for mamba/attn/shared_experts; expert (below) for routed + # experts captured when "moe" is in cuda_graph_modules. + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat", "expt_gtp_remat"] + ) + self._register_gtp_side_streams(pg_collection.gtp_remat) + # EGTP_remat streams: required so _wait/_sync_side_streams drain EGTP_remat + # NCCL into runner_stream before bwd_completion_event fires. + egtp_remat_group = pg_collection.expt_gtp_remat + if egtp_remat_group is not None and egtp_remat_group.size() > 1: + self._register_gtp_side_streams(egtp_remat_group) + # Registered for finalize_model_grads to wait on (Phase 2 fence). + _GTP_RUNNER_STREAMS.append(self.stream) if self.fp8_enabled: self.fp8_recipe = FP8GlobalStateManager.get_fp8_recipe() @@ -748,6 +985,56 @@ def __init__( self.fp4_recipe = get_fp4_recipe(self.base_module.config) _set_skip_fp8_weight_update_tensor(False) + def _register_gtp_side_streams(self, group): + """Register a GTP (chain, group)'s GRAPHED AG/RS side streams for capture/replay sync: the + AG stream on both fwd and bwd, the RS stream on bwd only.""" + ag = get_ag_stream(GTPChain.GRAPHED.value, group) + rs = get_rs_stream(GTPChain.GRAPHED.value, group) + self.fwd_side_streams.append(ag) + self.bwd_side_streams.append(ag) + self.bwd_side_streams.append(rs) + + def _sync_against_side_streams(self, side_streams): + """Make registered side streams wait for the current stream. + Also injects a dummy kernel into each stream to ensure it is non-empty, + which is required for CUDA graph capture (joining an empty captured + stream is a CUDA error).""" + for s in side_streams: + s.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(s): + torch.cuda._sleep(1) + + def _wait_side_streams(self, side_streams): + """Make the current stream wait for all registered side streams.""" + for s in side_streams: + torch.cuda.current_stream().wait_stream(s) + + def _compute_finalized_during_bwd_capture(self): + """Return GTP params whose DDP grad-ready hook fires post-replay + of THIS bwd_graph. + + A param's hook must fire in the graph that physically populates its + main_grad. Rules, given the cascade walk in wgrad_reduce_scatter + finalizes p.next_w on behalf of p: + - p.prev_w is None → p is sync-finalized in p's own graph; add p. + - p.next_w is not None → p.next_w's main_grad.add_ is captured here + via p's cascade; add p.next_w. (For cross-graph chain tails the + wait was captured in the producer's Phase 2, but the add lives + here regardless, bridged by external rs_event.) + """ + finalized = {} # id → param + for p in self.params_to_backprop: + if not getattr(p, 'is_gtp_weight_remat', False): + continue + if getattr(p, "prev_w", None) is None: + for w in getattr(p, "_weights", [p]): + finalized[id(w)] = w + next_w = getattr(p, "next_w", None) + if next_w is not None: + for w in getattr(next_w, "_weights", [next_w]): + finalized[id(w)] = w + return list(finalized.values()) + def __str__(self): return "%s; hid %s" % ( self.base_module.__class__.__name__, @@ -791,6 +1078,64 @@ def get_connected_params(self, outputs): # Return module params that were found in the graph, preserving original order return tuple(p for p in self.base_module.parameters() if id(p) in p_ids) + def _weakref_forward_buffers(self, preserve_forward_to_backward_lifetimes: bool) -> None: + """Release ownership only when CUDA graph topology proves the buffer reclaimable. + + `make_weakref` preserves a captured address but releases allocator ownership. + Although CUDA graph memory is pinned to a stable address, the graph-pool allocator may + reuse an unowned allocation before backward capture and overwrite its contents. + these conditions only within that interval avoids retaining every boundary tensor. + """ + + def is_saved_for_backward(tensor) -> bool: + """Return whether a tensor is needed for the backward pass graph. + + Preserving allocator ownership guards against the graph pool reusing and overwriting + that storage before backward capture records the read. + """ + + metadata = getattr(tensor, "cg_buffer_metadata", None) + return bool( + torch.is_tensor(tensor) and metadata is not None and metadata.is_saved_for_backward + ) + + def is_differentiable_cudagraph_output_escape(tensor) -> bool: + """Return whether a differentiable graph output escapes to eager code. + + Outputs that are also inputs to another CUDA graph are protected by graph-to-graph + reuse accounting. However, an output that is not another graph's input has no + such owner. Preserving it's ownership guards against premature graph-pool storage + reuse across that graph boundary. + """ + metadata = getattr(tensor, "cg_buffer_metadata", None) + return bool( + torch.is_tensor(tensor) + and tensor.requires_grad + and metadata is not None + and metadata.is_cudagraph_output + and not metadata.is_cudagraph_input + ) + + def weakref_input(tensor): + if preserve_forward_to_backward_lifetimes: + if is_saved_for_backward(tensor): + return tensor + return make_weakref(tensor) + + def weakref_output(tensor): + if preserve_forward_to_backward_lifetimes: + if is_saved_for_backward(tensor): + return tensor + if is_differentiable_cudagraph_output_escape(tensor): + return tensor + return make_weakref(tensor) + + self.fwd_graph_input_surface = tree_map(weakref_input, self.fwd_graph_input_surface) + self.fwd_graph_input_args = tree_map(weakref_input, self.fwd_graph_input_args) + self.fwd_graph_input_kwargs = tree_map(weakref_input, self.fwd_graph_input_kwargs) + self.fwd_graph_outputs = tree_map(weakref_output, self.fwd_graph_outputs) + self.fwd_graph_output_surface = tree_map(weakref_output, self.fwd_graph_output_surface) + def create_fwd_graph(self, args, kwargs, outputs=None, clone_inputs=True): """Create a fwd cudagraph for this runner. Should be called inside 'create_cudagraphs()'.""" @@ -813,9 +1158,7 @@ def create_fwd_graph(self, args, kwargs, outputs=None, clone_inputs=True): for buf in self.base_module.buffers(): buffer_backup.append(buf.clone()) - grad_backup = [] - for param in self.base_module.parameters(): - grad_backup.append(param.main_grad.clone() if hasattr(param, "main_grad") else None) + grad_backup = _backup_grads_before_capture(self) saved_fp8_tensors = None if self.fp8_enabled: @@ -858,52 +1201,49 @@ def create_fwd_graph(self, args, kwargs, outputs=None, clone_inputs=True): def _resolve_input_buffer(ten): if not isinstance(ten, ArgMetadata): return ten + metadata = getattr(ten, "cg_buffer_metadata", None) + # the input tensor is resued from another cudagraph's input or output - if ( - hasattr(ten, "cg_buffer_metadata") - and ten.cg_buffer_metadata.fwd_cudagraph_buffer is not None - ): - buf = ten.cg_buffer_metadata.fwd_cudagraph_buffer + if metadata is not None and metadata.fwd_cudagraph_buffer is not None: + shared_buf = metadata.fwd_cudagraph_buffer + buf_metadata = shared_buf.cg_buffer_metadata - assert ( - ten.cg_buffer_metadata.is_cudagraph_input - and buf.cg_buffer_metadata.capture_reuse_count > 0 - ) + assert metadata.is_cudagraph_input and buf_metadata.capture_reuse_count > 0 - if ( - ten.cg_buffer_metadata.input_use_count > 1 - and ten.cg_buffer_metadata.input_use_count - == buf.cg_buffer_metadata.capture_reuse_count - ): - can_skip_replay_copy = False - else: - can_skip_replay_copy = True + can_skip_replay_copy = not ( + metadata.input_use_count > 1 + and metadata.input_use_count == buf_metadata.capture_reuse_count + ) - buf.cg_buffer_metadata.capture_reuse_count -= 1 - if buf.cg_buffer_metadata.capture_reuse_count == 0: + buf_metadata.capture_reuse_count -= 1 + if buf_metadata.capture_reuse_count == 0: args_to_clear_buffers.append(ten) + + buf = create_strong_ref(shared_buf) else: # need to provide a fresh buffer from the pool buf = alloc_tensor_from_graph_mempool(ten) + if metadata is not None: + buf.cg_buffer_metadata = deepcopy(metadata) can_skip_replay_copy = False buf.can_skip_replay_copy = can_skip_replay_copy return buf if clone_inputs: - # if a buffer is used for multiple inputs, create it now - for ten in self.get_tensors(args, kwargs): + # Recorded graph arguments are ArgMetadata, not tensors. Preallocate a shared + # buffer before resolving each occurrence so later graph inputs can alias it. + for ten in self.get_arg_metas(args, kwargs): + metadata = getattr(ten, "cg_buffer_metadata", None) if ( - hasattr(ten, 'cg_buffer_metadata') - and ten.cg_buffer_metadata.input_use_count > 1 - and ten.cg_buffer_metadata.fwd_cudagraph_buffer is None + metadata is not None + and metadata.input_use_count > 1 + and metadata.fwd_cudagraph_buffer is None ): buf = alloc_tensor_from_graph_mempool(ten) - buf.cg_buffer_metadata = deepcopy(ten.cg_buffer_metadata) - buf.cg_buffer_metadata.capture_reuse_count = ( - ten.cg_buffer_metadata.input_use_count - ) - ten.cg_buffer_metadata.fwd_cudagraph_buffer = buf + buf.cg_buffer_metadata = deepcopy(metadata) + buf.cg_buffer_metadata.capture_reuse_count = metadata.input_use_count + metadata.fwd_cudagraph_buffer = buf fwd_buffer_reuse_ref_count += 1 self.fwd_graph_input_args = tree_map(_resolve_input_buffer, args) @@ -943,6 +1283,10 @@ def clone_ten(ten): allow_unused=True, ) + if self.gtp_remat: + wait_async_comms(GTPChain.GRAPHED.value) + self._sync_against_side_streams(self.bwd_side_streams) + _set_warmup_end() with self.get_quantization_context(): @@ -963,10 +1307,23 @@ def clone_ten(ten): with torch.cuda.graph( self.fwd_graph, pool=self.mempool, capture_error_mode="thread_local" ): + + self._sync_against_side_streams(self.fwd_side_streams) + fwd_graph_outputs = self.func( *self.fwd_graph_input_args, **self.fwd_graph_input_kwargs ) + if self.gtp_remat: + # Forward only issues AG prefetches (no wgrad RS), so drain AG and skip RS. + wait_async_comms(GTPChain.GRAPHED.value, skip_rs=True) + + if self.fwd_side_streams: + self._wait_side_streams(self.fwd_side_streams) + + if self.use_stream: + self.fwd_completion_event.record() + # Unfreeze GC. if FREEZE_GC: gc.unfreeze() @@ -990,19 +1347,15 @@ def clone_ten(ten): for fwd_graph_out, o in zip( self.get_tensors(fwd_graph_outputs), self.get_arg_metas(self.outputs) ): - assert hasattr(o, "cg_buffer_metadata") and o.cg_buffer_metadata.is_cudagraph_output + metadata = getattr(o, "cg_buffer_metadata", None) + assert metadata is not None and metadata.is_cudagraph_output fwd_graph_out.is_from_global_mempool = True - fwd_graph_out.cg_buffer_metadata = deepcopy(o.cg_buffer_metadata) + fwd_graph_out.cg_buffer_metadata = deepcopy(metadata) - if ( - o.cg_buffer_metadata.is_cudagraph_input - and o.cg_buffer_metadata.fwd_cudagraph_buffer is None - ): + if metadata.is_cudagraph_input and metadata.fwd_cudagraph_buffer is None: buf = create_strong_ref(fwd_graph_out) - buf.cg_buffer_metadata.capture_reuse_count = ( - o.cg_buffer_metadata.cudagraph_reuse_ref_count - ) - o.cg_buffer_metadata.fwd_cudagraph_buffer = buf + buf.cg_buffer_metadata.capture_reuse_count = metadata.cudagraph_reuse_ref_count + metadata.fwd_cudagraph_buffer = buf fwd_buffer_reuse_ref_count += 1 if self.training and torch.is_grad_enabled(): @@ -1012,11 +1365,8 @@ def clone_ten(ten): however the graphed module must output at least one tensor, so that a corresponding backward node may be registered in the autograd graph.""" - self.fwd_graph_input_surface = tree_map(make_weakref, self.fwd_graph_input_surface) - self.fwd_graph_input_args = tree_map(make_weakref, self.fwd_graph_input_args) - self.fwd_graph_input_kwargs = tree_map(make_weakref, self.fwd_graph_input_kwargs) - self.fwd_graph_outputs = tree_map(make_weakref, self.fwd_graph_outputs) - self.fwd_graph_output_surface = tree_map(make_weakref, self.fwd_graph_output_surface) + # Preserve only forward buffers whose lifetime crosses into backward capture. + self._weakref_forward_buffers(preserve_forward_to_backward_lifetimes=True) self.params_to_backprop = self.get_connected_params(fwd_graph_outputs) self.num_dgrads = len(self.fwd_graph_input_surface) @@ -1025,9 +1375,7 @@ def clone_ten(ten): if self.fp8_enabled: restore_fp8_tensors([self.base_module], saved_fp8_tensors) # restore cached grads - for main_grad_copy, param in zip(grad_backup, self.base_module.parameters()): - if main_grad_copy is not None: - param.main_grad.copy_(main_grad_copy) + _restore_grads_after_capture(grad_backup) # restore cached buffers for buf_copy, buf in zip(buffer_backup, self.base_module.buffers()): @@ -1060,17 +1408,15 @@ def create_bwd_graph(self): for o in self.get_arg_metas(self.outputs): out_grad = None if o.requires_grad: + metadata = o.cg_buffer_metadata # TODO: (jiemingz) [interaction with recompute] # for activation recompute, the fwd pass is rerun in the backward pass and # the metadata we attach in record_graph_capture is lost. As a result the next # cudagraph expects the buffer to be provided 'fwd_cudagraph_buffer' but is missing. # So, we cannot always assume this metadata exists. Consequently, there are extra # copies between the outputs of the fwd-bwd pass and the bwd pass. - if ( - o.cg_buffer_metadata.is_cudagraph_input - and o.cg_buffer_metadata.bwd_cudagraph_buffer is not None - ): - out_grad = o.cg_buffer_metadata.bwd_cudagraph_buffer + if metadata.is_cudagraph_input and metadata.bwd_cudagraph_buffer is not None: + out_grad = metadata.bwd_cudagraph_buffer args_to_clear_buffers.append(o) out_grad.cg_buffer_metadata.capture_reuse_count -= 1 else: @@ -1082,6 +1428,9 @@ def create_bwd_graph(self): gc.freeze() with torch.cuda.graph(self.bwd_graph, pool=self.mempool): + + self._sync_against_side_streams(self.bwd_side_streams) + grad_inputs = torch.autograd.grad( outputs=tuple(o for o in self.fwd_graph_output_surface if o.requires_grad), inputs=tuple(i for i in self.fwd_graph_input_surface if i.requires_grad), @@ -1098,10 +1447,71 @@ def create_bwd_graph(self): if wgrad is not None and not getattr(param, 'grad_added_to_main_grad', False): param.main_grad.add_(wgrad) + # GTP cross-graph RS overlap, two phases: + # Phase 1 — drain AG, fence runner_stream past ag_stream's tail, + # then record bwd_completion_event so main_stream can + # release the next runner while RS is still in flight. + # Phase 2 — drain RS wait on rs_stream. For cross-graph chain + # tails the wait is captured here, the add in the + # consumer's cascade; for within-graph tails both + # happen here (see wait_async_comms). + if self.gtp_remat: + # Phase 1: drain AG; fence runner_stream past dense + EGTP AG + # so bwd_completion_event records AFTER NCCL_AG completion. + wait_async_comms(GTPChain.GRAPHED.value, skip_rs=True) + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat", "expt_gtp_remat"] + ) + gtp_remat_group = pg_collection.gtp_remat + graphed_ag = get_ag_stream(GTPChain.GRAPHED.value, gtp_remat_group) + torch.cuda.current_stream().wait_stream(graphed_ag) + egtp_remat_group = pg_collection.expt_gtp_remat + if egtp_remat_group is not None and egtp_remat_group.size() > 1: + egtp_graphed_ag = get_ag_stream(GTPChain.GRAPHED.value, egtp_remat_group) + torch.cuda.current_stream().wait_stream(egtp_graphed_ag) + + # Record completion AFTER AG drain + fence but BEFORE RS drain, + # so main_stream can trigger the next runner while RS is still + # in flight on rs_stream. + self.bwd_completion_event.record() + + # Phase 2: in-graph RS drain + finalize. + wait_async_comms(GTPChain.GRAPHED.value, finalize_after_drain=True) + + if self.bwd_side_streams: + self._wait_side_streams(self.bwd_side_streams) + + if self.use_stream and not self.gtp_remat: + # Non-GTP path: record after the side-stream join. + self.bwd_completion_event.record() + # Unfreeze GC. if FREEZE_GC: gc.unfreeze() + # See _compute_finalized_during_bwd_capture for what's in this set and why. + self.finalized_during_bwd_capture = ( + self._compute_finalized_during_bwd_capture() if self.gtp_remat else [] + ) + + # Precompute the (rs_stream, params) DDP grad-ready hook plan once — it's + # replay-invariant — so Graphed.backward avoids per-replay group lookups. + self._gtp_finalize_hook_plan = [] + if self.gtp_remat and self.finalized_during_bwd_capture: + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat", "expt_gtp_remat"] + ) + dense_group = pg_collection.gtp_remat + expert_group = pg_collection.expt_gtp_remat + params_by_group = defaultdict(list) + for param in self.finalized_during_bwd_capture: + is_expert = not getattr(param, 'allreduce', True) + params_by_group[expert_group if is_expert else dense_group].append(param) + self._gtp_finalize_hook_plan = [ + (get_rs_stream(GTPChain.GRAPHED.value, group), params) + for group, params in params_by_group.items() + ] + for arg in args_to_clear_buffers: arg.cg_buffer_metadata.bwd_cudagraph_buffer = None bwd_buffer_reuse_ref_count -= 1 @@ -1113,16 +1523,14 @@ def create_bwd_graph(self): self.static_grad_inputs = [] for input_tensor in self.get_arg_metas(self.args, self.kwargs): if input_tensor.requires_grad: + metadata = input_tensor.cg_buffer_metadata input_grad = grad_inputs.pop(0) input_grad.is_from_global_mempool = True - input_grad.cg_buffer_metadata = deepcopy(input_tensor.cg_buffer_metadata) + input_grad.cg_buffer_metadata = deepcopy(metadata) - if ( - input_tensor.cg_buffer_metadata.is_cudagraph_output - and input_tensor.cg_buffer_metadata.bwd_cudagraph_buffer is None - ): + if metadata.is_cudagraph_output and metadata.bwd_cudagraph_buffer is None: buf = create_strong_ref(input_grad) - input_tensor.cg_buffer_metadata.bwd_cudagraph_buffer = buf + metadata.bwd_cudagraph_buffer = buf buf.cg_buffer_metadata.capture_reuse_count += 1 bwd_buffer_reuse_ref_count += 1 self.static_grad_inputs.append(input_grad) @@ -1137,6 +1545,8 @@ def create_bwd_graph(self): # stored in 'bwd_cudagraph_buffer' self.static_grad_inputs = tree_map(make_weakref, self.static_grad_inputs) self.static_grad_outputs = tree_map(make_weakref, self.static_grad_outputs) + # Backward capture is the final recorded use of forward buffers retained for autograd. + self._weakref_forward_buffers(preserve_forward_to_backward_lifetimes=False) delattr(self, "args") delattr(self, "kwargs") @@ -1146,19 +1556,15 @@ def apply_cudagraph_record_metadata(self, args, kwargs, outputs): """Attaches graph capture metadata to all passed in tensors.""" for t in self.get_tensors(args, kwargs): - if not hasattr(t, "cg_buffer_metadata"): - t.cg_buffer_metadata = CudagraphBufferMetadata() - - t.cg_buffer_metadata.is_cudagraph_input = True - t.cg_buffer_metadata.input_use_count += 1 + cg_buffer_metadata = _apply_cudagraph_buffer_metadata(t) + cg_buffer_metadata.is_cudagraph_input = True + cg_buffer_metadata.input_use_count += 1 - if t.cg_buffer_metadata.is_cudagraph_output: - t.cg_buffer_metadata.cudagraph_reuse_ref_count += 1 + if cg_buffer_metadata.is_cudagraph_output: + cg_buffer_metadata.cudagraph_reuse_ref_count += 1 - # mark all outputs, so that the fwd graph we may reuse cudagraph output buffers as inputs - for o in self.get_tensors(outputs): - o.cg_buffer_metadata = CudagraphBufferMetadata() - o.cg_buffer_metadata.is_cudagraph_output = True + for t in self.get_tensors(outputs): + _apply_cudagraph_buffer_metadata(t, is_output=True) def record_graph_capture(self, args, kwargs): """Records the data needed to create this runner's forward cudagraph. @@ -1432,10 +1838,19 @@ def wrapped_func(*args, eager=False, cache_key=None, **kwargs): self.reuse_cudagraphs = self.pg_collection.pp.size() == 1 if CudaGraphManager.global_mempool is None: CudaGraphManager.global_mempool = torch.cuda.graph_pool_handle() + # Register the pool so GTP allocates GRAPHED-chain buffers + quantized + # storage directly into it (created before the first graphed forward). + if HAVE_GTP: + set_cuda_graph_mempool(torch.cuda.current_device(), CudaGraphManager.global_mempool) # Cudagraph stream capture requires no operations on the default stream prior to the # capture, so change to a side stream. torch.cuda.set_stream(torch.cuda.Stream()) + # Enable one hook for the eager recording phase. Repeated manager construction is + # idempotent, and graph creation removes the hook before capture begins. + if need_backward: + _CudagraphGlobalRecord._enable_saved_tensors_observer() + def call_ddp_preforward_hook(self, module): """Call any DDP pre-forward hooks which are used to launch async data parallel param gather. Any other pre-forward hooks are not allowed.""" @@ -1618,7 +2033,7 @@ def __call__(self, megatron_module, args, kwargs, cache_key=None): self.is_first_microbatch = False # If forward only, next replay should be a forward pass as well - if is_inference_mode or not torch.is_grad_enabled(): + if is_inference_mode or not torch.is_grad_enabled() or not runner.fwd_graph_recorded: runner.status = _GraphStatus.FWD_READY else: runner.status = _GraphStatus.BWD_READY diff --git a/megatron/core/transformer/custom_layers/batch_invariant_kernels.py b/megatron/core/transformer/custom_layers/batch_invariant_kernels.py index 6b4311fe540..a83e298eee2 100644 --- a/megatron/core/transformer/custom_layers/batch_invariant_kernels.py +++ b/megatron/core/transformer/custom_layers/batch_invariant_kernels.py @@ -525,10 +525,15 @@ def get_batch_invariant_attention_block_size() -> AttentionBlockSize: _MEG_TE_GENERAL_GEMM_ORIG = None _TE_RMSNORM_FUNC_ORIGS: Dict[str, Any] = {} _TE_GEMM_FUNC_ORIGS: Dict[str, Any] = {} +_TE_APPLY_NORM_ORIGS: Dict[str, Any] = {} def _import_module_if_available(name: str): - spec = importlib.util.find_spec(name) + try: + spec = importlib.util.find_spec(name) + except ModuleNotFoundError: + # find_spec on a submodule raises when the parent package is absent. + return None if spec is None: return None return importlib.import_module(name) @@ -617,6 +622,34 @@ def _patched(*args, **kwargs): _TE_RMSNORM_FUNC_ORIGS[name] = orig setattr(te_layernorm_mod, name, _make_rmsnorm_patched(orig)) + # Patch the fused-module normalization entry (`apply_normalization`). TE's + # fused LayerNormLinear / LayerNormMLP call this instead of RMSNorm.forward, + # so without this patch their internal RMSNorm runs TE's tex kernel, whose + # within-row reduction strategy depends on the total row count — i.e. it is + # NOT batch-invariant (observed: same rows, different output at 928 vs 2274 + # rows on GB200, 1 bf16 ulp per layer, amplifying across depth). + import transformer_engine.pytorch.module._common as te_common + + for mod_name, mod in ( + ("module._common", te_common), + ( + "module.layernorm_linear", + _import_module_if_available("transformer_engine.pytorch.module.layernorm_linear"), + ), + ( + "module.layernorm_mlp", + _import_module_if_available("transformer_engine.pytorch.module.layernorm_mlp"), + ), + ): + key = f"{mod_name}.apply_normalization" + if ( + mod is not None + and hasattr(mod, "apply_normalization") + and key not in _TE_APPLY_NORM_ORIGS + ): + _TE_APPLY_NORM_ORIGS[key] = mod.apply_normalization + mod.apply_normalization = _te_apply_normalization_patched + def _te_unpatch_for_batch_invariant(): """Restore original Transformer Engine functions if they were patched.""" @@ -652,6 +685,14 @@ def _te_unpatch_for_batch_invariant(): elif meg_te is None: _MEG_TE_GENERAL_GEMM_ORIG = None + # Restore fused-module apply_normalization entries + for key, orig in list(_TE_APPLY_NORM_ORIGS.items()): + mod_name = key.rsplit(".apply_normalization", 1)[0] + mod = _import_module_if_available(f"transformer_engine.pytorch.{mod_name}") + if mod is not None and hasattr(mod, "apply_normalization"): + mod.apply_normalization = orig + _TE_APPLY_NORM_ORIGS.pop(key, None) + # Restore TE module-level RMSNorm functions te_layernorm_mod = _import_module_if_available("transformer_engine.pytorch.module.layernorm") if te_layernorm_mod is not None: @@ -827,6 +868,72 @@ def backward(ctx, grad_output: torch.Tensor): return dA, dB, dbias, None, None +def _te_apply_normalization_patched( + inputmat, + ln_out, + ln_weight, + ln_bias, + eps, + output_quantizer, + output_dtype, + normalization, + fwd_ln_sm_margin, + zero_centered_gamma, +): + """Batch-invariant replacement for TE's fused-module `apply_normalization`. + + Routes RMSNorm through the batch-invariant implementation (fp32 stats via + `mean_dim`, deterministic within-row reduction independent of row count). + Falls back to the original TE kernel for configurations the BI path does + not cover (LayerNorm, fp8 quantized output). + + Returns `(ln_out, mu, rsigma)` like TE: `mu` is None for RMSNorm and + `rsigma` is fp32 with shape [rows], which TE's rmsnorm backward consumes. + """ + orig = _TE_APPLY_NORM_ORIGS.get("module._common.apply_normalization") + if ( + not is_batch_invariant_mode_enabled() + or normalization != "RMSNorm" + or output_quantizer is not None + or ln_bias is not None + ): + assert orig is not None, "TE apply_normalization original not captured" + return orig( + inputmat, + ln_out, + ln_weight, + ln_bias, + eps, + output_quantizer, + output_dtype, + normalization, + fwd_ln_sm_margin, + zero_centered_gamma, + ) + + x_fp32 = inputmat.float() + w_fp32 = ln_weight.float() + if zero_centered_gamma: + w_fp32 = w_fp32 + 1.0 + ms = mean_dim(x_fp32 * x_fp32, dim=-1, keepdim=True) + rsigma = torch.rsqrt(ms + eps) + out_fp32 = (x_fp32 * rsigma) * w_fp32 + + # The fused-module callers (layernorm_linear / layernorm_mlp) pass a torch.dtype + # output_dtype (inputmat.dtype); assert that so we never silently ignore a TE + # DType or other unexpected value here. + assert isinstance(output_dtype, torch.dtype), ( + "batch-invariant apply_normalization expects a torch.dtype output_dtype, got " + f"{type(output_dtype)}" + ) + if ln_out is not None: + # copy_ casts to ln_out.dtype in place, avoiding an intermediate allocation. + ln_out.copy_(out_fp32) + else: + ln_out = out_fp32.to(output_dtype) + return ln_out, None, rsigma.squeeze(-1) + + def _te_general_gemm_patched(*args, **kwargs) -> List[torch.Tensor]: """ Batch-invariant replacement for TE general_gemm. diff --git a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py index 64a08cb7837..8ecbf70045a 100644 --- a/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py +++ b/megatron/core/transformer/experimental_attention_variant/absorbed_mla.py @@ -36,6 +36,7 @@ ) from megatron.core.transformer.attention import Attention from megatron.core.transformer.enums import AttnMaskType +from megatron.core.transformer.mla_qk_norm_config import QKNormConfigResolver from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.transformer_config import MLATransformerConfig from megatron.core.utils import deprecate_inference_params, get_pg_size, is_te_min_version @@ -170,6 +171,10 @@ def __init__( name=name, ) + # Resolve which classes to use for Q and KV linear up projections and norms, based on + # QK-norm selection. + layer_classes = QKNormConfigResolver(self.config, submodules).resolve() + assert not config.add_bias_linear, "add_bias_linear is not supported for AbsorbedMLA" assert not ( config.tensor_model_parallel_size > 1 and not config.sequence_parallel @@ -269,7 +274,7 @@ def __init__( if self.config.q_lora_rank is None: # Not projecting query self.linear_q_proj = build_module( - submodules.linear_q_proj, + layer_classes["linear_q_proj"], self.config.hidden_size, self.config.num_attention_heads * self.q_head_dim, config=self.config, @@ -315,7 +320,7 @@ def __init__( ) self.linear_q_up_proj = build_module( - submodules.linear_q_up_proj, + layer_classes["linear_q_up_proj"], self.config.q_lora_rank, self.config.num_attention_heads * self.q_head_dim, config=self.config, @@ -422,14 +427,14 @@ def __init__( if self.config.q_lora_rank is not None: self.q_layernorm = build_module( - submodules.q_layernorm, + layer_classes["q_layernorm"], hidden_size=self.config.q_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, ) self.kv_layernorm = build_module( - submodules.kv_layernorm, + layer_classes["kv_layernorm"], hidden_size=self.config.kv_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, diff --git a/megatron/core/transformer/experimental_attention_variant/dsa_cudnn_kernels.py b/megatron/core/transformer/experimental_attention_variant/dsa_cudnn_kernels.py index e26c46030aa..640db5c91ec 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa_cudnn_kernels.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa_cudnn_kernels.py @@ -725,42 +725,15 @@ def _indexer_topk_multi_packed_cp_thd( raise RuntimeError("packed CP cuDNN THD indexer requires positive maximum sequence lengths") segment_divisor = 2 * cp_size - if sk % segment_divisor != 0: - raise RuntimeError(f"packed CP key length must be divisible by {segment_divisor}, got {sk}") - device = q_bshd.device - cu_q = packed_cu_seqlens_q.to(device=device, dtype=torch.int64).contiguous() - cu_k = packed_cu_seqlens_k.to(device=device, dtype=torch.int64).contiguous() - q_lengths = cu_q[1:] - cu_q[:-1] - k_lengths = cu_k[1:] - cu_k[:-1] - q_half = q_lengths // segment_divisor - k_half = k_lengths // segment_divisor - segment_q_lengths = torch.stack((q_half, q_half), dim=1).reshape(-1) - segment_k_lengths = torch.stack( - ((cp_rank + 1) * k_half, k_lengths - cp_rank * k_half), dim=1 - ).reshape(-1) - - zero_i32 = torch.zeros(1, dtype=torch.int32, device=device) - segment_cu_q = torch.cat( - (zero_i32, segment_q_lengths.cumsum(dim=0, dtype=torch.int32)) - ).contiguous() - segment_cu_k = torch.cat( - (zero_i32, segment_k_lengths.cumsum(dim=0, dtype=torch.int32)) - ).contiguous() - - segment_key_starts = cu_k[:-1].repeat_interleave(2) - total_segment_k = sk + sk // segment_divisor - segment_ids = torch.repeat_interleave( - torch.arange(segment_k_lengths.numel(), device=device), - segment_k_lengths, - output_size=total_segment_k, - ) - segment_offsets = torch.arange(total_segment_k, device=device, dtype=torch.int64) - segment_offsets -= torch.repeat_interleave( - segment_cu_k[:-1].to(dtype=torch.int64), segment_k_lengths, output_size=total_segment_k + layout = dsa_layout.build_packed_cp_indexer_layout( + packed_cu_seqlens_q.to(device=device), + packed_cu_seqlens_k.to(device=device), + cp_size=cp_size, + cp_rank=cp_rank, + key_size=sk, ) - source_indices = segment_key_starts.index_select(0, segment_ids) + segment_offsets - segmented_k = k_bshd[0].index_select(0, source_indices).contiguous() + segmented_k = k_bshd[0].index_select(0, layout.source_indices).contiguous() max_segment_q = packed_max_seqlen_q // segment_divisor max_k_half = packed_max_seqlen_k // segment_divisor @@ -771,8 +744,8 @@ def _indexer_topk_multi_packed_cp_thd( w_bsh[0], ratio=_INDEXER_RATIO, sm_scale=_INDEXER_SOFTMAX_SCALE, - cu_seqlens_q=segment_cu_q, - cu_seqlens_k=segment_cu_k, + cu_seqlens_q=layout.segment_cu_q.to(dtype=torch.int32), + cu_seqlens_k=layout.segment_cu_k.to(dtype=torch.int32), max_seqlen_q=max_segment_q, max_seqlen_k=max_segment_k, )["scores"] @@ -1043,11 +1016,8 @@ def _sort_valid_topk_indices_by_index(topk_indices: Tensor, topk_length: Tensor, """Canonicalize consumed top-K indices while keeping ignored suffix slots invalid.""" positions = _trailing_positions(topk_indices) valid = positions < topk_length.unsqueeze(-1) - sort_key = torch.where(valid, topk_indices, torch.full_like(topk_indices, sk)) - order = sort_key.argsort(dim=-1) - sorted_indices = torch.gather(topk_indices, dim=-1, index=order) - sorted_valid = torch.gather(valid.expand_as(topk_indices), dim=-1, index=order) - return sorted_indices.masked_fill(~sorted_valid, -1).contiguous() + sorted_indices, _ = dsa_masking.sort_topk_by_index(topk_indices, valid, sk=sk) + return sorted_indices def _sort_valid_topk_indices_and_scores_by_index( @@ -1056,14 +1026,15 @@ def _sort_valid_topk_indices_and_scores_by_index( """Sort valid top-K indices and keep the selected score payload aligned.""" positions = _trailing_positions(topk_indices) valid = positions < topk_length.unsqueeze(-1) - sort_key = torch.where(valid, topk_indices, torch.full_like(topk_indices, sk)) - order = sort_key.argsort(dim=-1) - sorted_indices = torch.gather(topk_indices, dim=-1, index=order) - sorted_scores = torch.gather(topk_scores, dim=-1, index=order) - sorted_valid = torch.gather(valid.expand_as(topk_indices), dim=-1, index=order) - sorted_indices = sorted_indices.masked_fill(~sorted_valid, -1) - sorted_scores = sorted_scores.masked_fill(~sorted_valid, torch.finfo(torch.float32).min) - return sorted_indices.contiguous(), sorted_scores.contiguous() + sorted_indices, sorted_scores = dsa_masking.sort_topk_by_index( + topk_indices, + valid, + sk=sk, + topk_scores=topk_scores, + invalid_score=torch.finfo(torch.float32).min, + ) + assert sorted_scores is not None + return sorted_indices, sorted_scores def _prepare_attention_topk_indices(topk_indices: Tensor, sk: int) -> Tuple[Tensor, Tensor]: diff --git a/megatron/core/transformer/experimental_attention_variant/dsa_layout.py b/megatron/core/transformer/experimental_attention_variant/dsa_layout.py index fd0ef34c701..4d4aa717b83 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa_layout.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa_layout.py @@ -2,6 +2,7 @@ """Layout helpers for DeepSeek sparse attention.""" +from dataclasses import dataclass from typing import Optional, Tuple import torch @@ -10,6 +11,8 @@ from megatron.core.utils import get_pg_size __all__ = [ + "PackedCPIndexerLayout", + "build_packed_cp_indexer_layout", "build_packed_allgather_cp_local_positions", "build_packed_allgather_cp_query_positions_and_key_reorder", "build_zigzag_allgather_cp_key_reorder", @@ -22,6 +25,93 @@ ] +@dataclass(frozen=True) +class PackedCPIndexerLayout: + """Segment metadata shared by packed-CP DSA indexer backends.""" + + segment_q_lengths: torch.Tensor + segment_k_lengths: torch.Tensor + segment_cu_q: torch.Tensor + segment_cu_k: torch.Tensor + segment_key_starts: torch.Tensor + source_indices: torch.Tensor + + +def build_packed_cp_indexer_layout( + cu_seqlens_q: torch.Tensor, + cu_seqlens_kv: torch.Tensor, + *, + cp_size: int, + cp_rank: int, + key_size: int, + local_key_layout: bool = False, +) -> PackedCPIndexerLayout: + """Build packed-CP front/back segment metadata for fused DSA indexers. + + ``local_key_layout`` describes the single-sequence optimization where the + key tensor contains only this CP rank's local front/back chunks. Otherwise, + ``key_size`` is the globally ordered packed key length. + """ + if cp_size <= 1 or not 0 <= cp_rank < cp_size: + raise RuntimeError("packed CP indexer layout requires a valid CP rank and cp_size > 1") + if cu_seqlens_q.shape != cu_seqlens_kv.shape or cu_seqlens_q.numel() < 2: + raise RuntimeError("packed CP indexer layout requires matching non-empty q/k cu_seqlens") + + device = cu_seqlens_q.device + cu_q = cu_seqlens_q.to(device=device, dtype=torch.int64).contiguous() + cu_k = cu_seqlens_kv.to(device=device, dtype=torch.int64).contiguous() + segment_divisor = 2 * cp_size + + if local_key_layout: + if cu_q.numel() != 2 or key_size % 2 != 0: + raise RuntimeError( + "local-key packed CP indexer layout requires one sequence and even key rows" + ) + half = key_size // 2 + segment_q_lengths = torch.full((2,), half, dtype=torch.int64, device=device) + segment_k_lengths = torch.tensor((half, key_size), dtype=torch.int64, device=device) + segment_key_starts = torch.zeros(2, dtype=torch.int64, device=device) + total_segment_k = key_size + half + else: + if key_size % segment_divisor != 0: + raise RuntimeError( + f"packed CP key length must be divisible by {segment_divisor}, got {key_size}" + ) + q_lengths = cu_q[1:] - cu_q[:-1] + k_lengths = cu_k[1:] - cu_k[:-1] + q_half = q_lengths // segment_divisor + k_half = k_lengths // segment_divisor + segment_q_lengths = torch.stack((q_half, q_half), dim=1).reshape(-1) + segment_k_lengths = torch.stack( + ((cp_rank + 1) * k_half, k_lengths - cp_rank * k_half), dim=1 + ).reshape(-1) + segment_key_starts = cu_k[:-1].repeat_interleave(2) + total_segment_k = key_size + key_size // segment_divisor + + zero = torch.zeros(1, dtype=torch.int64, device=device) + segment_cu_q = torch.cat((zero, segment_q_lengths.cumsum(dim=0))).contiguous() + segment_cu_k = torch.cat((zero, segment_k_lengths.cumsum(dim=0))).contiguous() + + segment_ids = torch.repeat_interleave( + torch.arange(segment_k_lengths.numel(), device=device), + segment_k_lengths, + output_size=total_segment_k, + ) + segment_offsets = torch.arange(total_segment_k, device=device, dtype=torch.int64) + segment_offsets -= torch.repeat_interleave( + segment_cu_k[:-1], segment_k_lengths, output_size=total_segment_k + ) + source_indices = segment_key_starts.index_select(0, segment_ids) + segment_offsets + return PackedCPIndexerLayout( + segment_q_lengths=segment_q_lengths, + segment_k_lengths=segment_k_lengths, + segment_cu_q=segment_cu_q, + segment_cu_k=segment_cu_k, + segment_key_starts=segment_key_starts, + source_indices=source_indices, + ) + + def normalize_cp_comm_type(cp_comm_type: Optional[str]) -> str: """Normalize CP communication type to a canonical lowercase form.""" if cp_comm_type is None: diff --git a/megatron/core/transformer/experimental_attention_variant/dsa_masking.py b/megatron/core/transformer/experimental_attention_variant/dsa_masking.py index 6bd1669098c..c66d066ac99 100644 --- a/megatron/core/transformer/experimental_attention_variant/dsa_masking.py +++ b/megatron/core/transformer/experimental_attention_variant/dsa_masking.py @@ -29,6 +29,7 @@ "prepare_additive_mask", "prepare_sparse_mask_context", "scatter_topk_into_index_mask", + "sort_topk_by_index", ] @@ -98,6 +99,37 @@ def build_valid_mask_from_starts_ends( ) +def sort_topk_by_index( + topk_indices: torch.Tensor, + valid_mask: torch.Tensor, + *, + sk: int, + topk_scores: Optional[torch.Tensor] = None, + invalid_score: float = float("-inf"), +) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + """Sort valid top-k slots by key index while preserving aligned scores. + + Backends define validity explicitly: TileLang uses ``index >= 0`` sentinels, + while cuDNN consumes a compact prefix described by ``topk_length``. + """ + if valid_mask.dtype != torch.bool or valid_mask.shape != topk_indices.shape: + raise ValueError("valid_mask must be boolean and match topk_indices") + if topk_scores is not None and topk_scores.shape != topk_indices.shape: + raise ValueError("topk_scores must match topk_indices") + + sort_key = torch.where(valid_mask, topk_indices, torch.full_like(topk_indices, sk)) + order = sort_key.argsort(dim=-1) + sorted_valid = torch.gather(valid_mask, dim=-1, index=order) + sorted_indices = torch.gather(topk_indices, dim=-1, index=order) + sorted_indices = sorted_indices.masked_fill(~sorted_valid, -1).contiguous() + if topk_scores is None: + return sorted_indices, None + + sorted_scores = torch.gather(topk_scores, dim=-1, index=order) + sorted_scores = sorted_scores.masked_fill(~sorted_valid, invalid_score).contiguous() + return sorted_indices, sorted_scores + + def apply_starts_ends_mask_to_scores( scores: torch.Tensor, starts: torch.Tensor, ends: torch.Tensor, key_positions: torch.Tensor ) -> torch.Tensor: diff --git a/megatron/core/transformer/experimental_attention_variant/dsa_tilelang_kernels.py b/megatron/core/transformer/experimental_attention_variant/dsa_tilelang_kernels.py new file mode 100644 index 00000000000..1d89e43e196 --- /dev/null +++ b/megatron/core/transformer/experimental_attention_variant/dsa_tilelang_kernels.py @@ -0,0 +1,142 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""TileLang backend hooks for optional fused DeepSeek sparse attention kernels.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Optional, Tuple + +import torch + +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.transformer.experimental_attention_variant.ops import tilelang_dsa + +if TYPE_CHECKING: + from megatron.core.packed_seq_params import PackedSeqParams + from megatron.core.transformer.transformer_config import TransformerConfig + + +def run_fused_qk_topk( + q: torch.Tensor, + k: torch.Tensor, + weights: torch.Tensor, + index_topk: int, + starts: torch.Tensor, + ends: torch.Tensor, + block_size: int, + use_relu: bool = True, + use_local_indexer_varlen: bool = False, + single_packed_thd_sequence: bool = False, + local_packed_cp_rank: int = 0, + local_packed_cp_query_start: int = 0, + local_packed_cp_query_len: Optional[int] = None, + packed_seq_params: Optional[PackedSeqParams] = None, + cp_size: int = 1, +) -> Optional[Tuple[torch.Tensor, Optional[torch.Tensor]]]: + """Adapt TileLang's indices-only result to the shared backend hook contract.""" + topk_indices = tilelang_dsa.run_fused_qk_topk( + q, + k, + weights, + index_topk, + starts, + ends, + block_size, + use_relu, + use_local_indexer_varlen=use_local_indexer_varlen, + single_packed_thd_sequence=single_packed_thd_sequence, + local_packed_cp_rank=local_packed_cp_rank, + local_packed_cp_query_start=local_packed_cp_query_start, + local_packed_cp_query_len=local_packed_cp_query_len, + packed_seq_params=packed_seq_params, + cp_size=cp_size, + ) + if topk_indices is None: + return None + return topk_indices, None + + +def run_fused_qk_topk_with_loss( + q: torch.Tensor, + k: torch.Tensor, + weights: torch.Tensor, + index_topk: int, + starts: torch.Tensor, + ends: torch.Tensor, + block_size: int, + query: torch.Tensor, + key: torch.Tensor, + softmax_scale: float, + loss_coeff: float, + pg_collection: ProcessGroupCollection, + query_valid_rows: Optional[torch.Tensor] = None, + calculate_per_token_loss: bool = False, + use_relu: bool = True, + config: Optional["TransformerConfig"] = None, + use_local_indexer_varlen: bool = False, + single_packed_thd_sequence: bool = False, + local_packed_cp_rank: int = 0, + local_packed_cp_query_start: int = 0, + local_packed_cp_query_len: Optional[int] = None, + packed_seq_params: Optional[PackedSeqParams] = None, + cp_size: int = 1, +) -> Optional[Tuple[torch.Tensor, Optional[torch.Tensor], torch.Tensor]]: + """Run fused TileLang indexer and sparse indexer loss.""" + del config + result = tilelang_dsa.run_fused_qk_topk_with_loss( + q=q, + k=k, + weights=weights, + index_topk=index_topk, + starts=starts, + ends=ends, + block_size=block_size, + query=query, + key=key, + softmax_scale=softmax_scale, + loss_coeff=loss_coeff, + pg_collection=pg_collection, + query_valid_rows=query_valid_rows, + calculate_per_token_loss=calculate_per_token_loss, + use_relu=use_relu, + use_local_indexer_varlen=use_local_indexer_varlen, + single_packed_thd_sequence=single_packed_thd_sequence, + local_packed_cp_rank=local_packed_cp_rank, + local_packed_cp_query_start=local_packed_cp_query_start, + local_packed_cp_query_len=local_packed_cp_query_len, + packed_seq_params=packed_seq_params, + cp_size=cp_size, + ) + if result is None: + return None + topk_indices, indexer_loss = result + return topk_indices, None, indexer_loss + + +def run_fused_absorbed_sparse_attention( + query: torch.Tensor, + key: torch.Tensor, + topk_indices: torch.Tensor, + softmax_scale: float, + v_channels: int, + topk_length: Optional[torch.Tensor] = None, +) -> Optional[torch.Tensor]: + """Run fused TileLang SparseMLA for absorbed DSA sparse attention.""" + if topk_length is not None: + if topk_indices.ndim != 3 or topk_length.shape != topk_indices.shape[:-1]: + return None + positions = torch.arange(topk_indices.size(-1), device=topk_indices.device) + valid = positions < topk_length.to(dtype=torch.int64, device=topk_indices.device).unsqueeze( + -1 + ) + topk_indices = topk_indices.masked_fill(~valid, -1) + return tilelang_dsa.run_fused_absorbed_sparse_attention( + query, key, topk_indices, softmax_scale, v_channels + ) + + +__all__ = [ + "run_fused_absorbed_sparse_attention", + "run_fused_qk_topk", + "run_fused_qk_topk_with_loss", +] diff --git a/megatron/core/transformer/experimental_attention_variant/ops/indexer.py b/megatron/core/transformer/experimental_attention_variant/ops/indexer.py new file mode 100644 index 00000000000..2f59feb0776 --- /dev/null +++ b/megatron/core/transformer/experimental_attention_variant/ops/indexer.py @@ -0,0 +1,131 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import torch + +from .tilelang_indexer_bwd import HAVE_TILELANG as HAVE_TILELANG_INDEXER_BWD +from .tilelang_indexer_bwd import indexer_bwd_interface +from .tilelang_indexer_fwd import HAVE_TILELANG as HAVE_TILELANG_INDEXER_FWD +from .tilelang_indexer_fwd import indexer_fwd_interface + +HAVE_TILELANG_INDEXER = HAVE_TILELANG_INDEXER_BWD and HAVE_TILELANG_INDEXER_FWD + + +def pytorch_extract_topk_scores(logits, topk_indices, dim=-1): + """Gather top-k logits and mask invalid (-1) entries with -inf.""" + if logits.size(dim) == 0: + return torch.full( + topk_indices.shape, float("-inf"), dtype=logits.dtype, device=logits.device + ) + valid_mask = (topk_indices >= 0) & (topk_indices < logits.size(dim)) + safe_indices = topk_indices.clamp(min=0, max=logits.size(dim) - 1).to(torch.int64) + scores = torch.gather(logits, dim=dim, index=safe_indices) + scores = torch.where(valid_mask, scores, float("-inf")) + return scores + + +def _select_topk_from_logits( + logits: torch.Tensor, topk: int, mask_invalid: bool = True +) -> tuple[torch.Tensor, torch.Tensor]: + """Select top-k scores and int32 indices from indexer logits.""" + effective_topk = min(topk, logits.size(-1)) + if effective_topk > 0: + topk_scores, topk_indices = torch.topk(logits, effective_topk, dim=-1, sorted=False) + topk_indices = topk_indices.to(torch.int32) + if mask_invalid: + topk_indices = topk_indices.masked_fill(topk_scores == -torch.inf, -1) + return topk_scores, topk_indices + + empty_shape = logits.shape[:-1] + (0,) + topk_scores = torch.empty(empty_shape, dtype=logits.dtype, device=logits.device) + topk_indices = torch.empty(empty_shape, dtype=torch.int32, device=logits.device) + return topk_scores, topk_indices + + +class IndexerFunction(torch.autograd.Function): # pragma: no cover + """Autograd wrapper for fused tilelang indexer forward/backward.""" + + @staticmethod + def forward( + ctx, + index_q: torch.Tensor, + index_k: torch.Tensor, + weights: torch.Tensor, + cu_seqlen_ks: torch.Tensor, + cu_seqlen_ke: torch.Tensor, + topk: int, + topk_indices: torch.Tensor | None = None, + use_relu: bool = True, + ): + """Run fused indexer forward and optionally select top-k indices.""" + logits = indexer_fwd_interface( + index_q, + index_k, + weights, + cu_seqlen_ks, + cu_seqlen_ke, + clean_logits=True, + use_relu=use_relu, + ) + if topk_indices is None: + index_score, topk_indices = _select_topk_from_logits(logits, topk) + else: + index_score = pytorch_extract_topk_scores(logits, topk_indices) + + ctx.save_for_backward(index_q, index_k, weights, cu_seqlen_ks, cu_seqlen_ke, topk_indices) + ctx.use_relu = use_relu + return index_score, topk_indices + + @staticmethod + def backward(ctx, grad_scores, grad_indices): + """Propagate gradients through fused indexer outputs.""" + index_q, index_k, weights, cu_seqlen_ks, cu_seqlen_ke, topk_indices = ctx.saved_tensors + grad_q, grad_w, grad_k = indexer_bwd_interface( + index_q, weights, index_k, topk_indices, grad_scores, use_relu=ctx.use_relu + ) + return grad_q, grad_k, grad_w, None, None, None, None, None + + +def lighting_indexer( # pragma: no cover + index_q: torch.Tensor, + index_k: torch.Tensor, + weights: torch.Tensor, + cu_seqlen_ks: torch.Tensor, + cu_seqlen_ke: torch.Tensor, + topk: int, + topk_indices: torch.Tensor | None = None, + use_relu: bool = True, +): + """Compute indexer top-k scores/indices via the custom autograd function.""" + return IndexerFunction.apply( + index_q, index_k, weights, cu_seqlen_ks, cu_seqlen_ke, topk, topk_indices, use_relu + ) + + +def lighting_indexer_indices( # pragma: no cover + index_q: torch.Tensor, + index_k: torch.Tensor, + weights: torch.Tensor, + cu_seqlen_ks: torch.Tensor, + cu_seqlen_ke: torch.Tensor, + topk: int, + use_relu: bool = True, +): + """Compute TileLang indexer top-k indices without score/autograd bookkeeping.""" + with torch.no_grad(): + logits = indexer_fwd_interface( + index_q, + index_k, + weights, + cu_seqlen_ks, + cu_seqlen_ke, + clean_logits=True, + use_relu=use_relu, + ) + _, topk_indices = _select_topk_from_logits(logits, topk, mask_invalid=False) + return topk_indices + + +if not HAVE_TILELANG_INDEXER: + IndexerFunction = None + lighting_indexer = None + lighting_indexer_indices = None diff --git a/megatron/core/transformer/experimental_attention_variant/ops/sparse_mla.py b/megatron/core/transformer/experimental_attention_variant/ops/sparse_mla.py new file mode 100644 index 00000000000..76976ebe0aa --- /dev/null +++ b/megatron/core/transformer/experimental_attention_variant/ops/sparse_mla.py @@ -0,0 +1,105 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import torch + +from .tilelang_sparse_mla_bwd import HAVE_TILELANG as HAVE_TILELANG_SPARSE_MLA_BWD +from .tilelang_sparse_mla_bwd import sparse_mla_bwd, sparse_mla_delta +from .tilelang_sparse_mla_fwd import HAVE_TILELANG as HAVE_TILELANG_SPARSE_MLA_FWD +from .tilelang_sparse_mla_fwd import sparse_mla_fwd_interface + +HAVE_TILELANG_SPARSE_MLA = HAVE_TILELANG_SPARSE_MLA_BWD and HAVE_TILELANG_SPARSE_MLA_FWD + + +def _canonicalize_batch_stride(tensor: torch.Tensor) -> torch.Tensor: + """Normalize a size-one batch stride without copying tensor data.""" + tensor = tensor.contiguous() + if tensor.ndim == 4 and tensor.size(0) == 1: + tensor = tensor.squeeze(0).unsqueeze(0) + return tensor + + +def _valid_head_mask(indices, num_heads): + valid_groups = indices.ge(0).any(dim=-1) + kv_group = valid_groups.size(-1) + if kv_group == num_heads: + return valid_groups + if num_heads % kv_group != 0: + raise RuntimeError( + f"SparseMLA heads must be divisible by kv_group, got heads={num_heads}, " + f"kv_group={kv_group}" + ) + return valid_groups.repeat_interleave(num_heads // kv_group, dim=-1) + + +def _zero_invalid_heads(tensor, valid_heads): + zero = torch.zeros((), dtype=tensor.dtype, device=tensor.device) + return torch.where(valid_heads.unsqueeze(-1), tensor, zero) + + +class SparseMLA(torch.autograd.Function): # pragma: no cover + """Autograd wrapper around tilelang sparse-MLA forward/backward kernels.""" + + @staticmethod + def forward(ctx, q, kv, indices, scaling): + """ + Args: + q: Query tensor (seq_len, heads, dim_plus_tail_dim) or + (batch, seq_len, heads, dim_plus_tail_dim) + kv: Key-Value tensor (seq_len_kv, kv_group, dim_plus_tail_dim) or + (batch, seq_len_kv, kv_group, dim_plus_tail_dim) + indices: Sparse indices tensor (seq_len, kv_group, topk) or + (batch, seq_len, kv_group, topk) + + Returns: + out: Output tensor (seq_len, heads, dim) or (batch, seq_len, heads, dim) + """ + indices = _canonicalize_batch_stride(indices) + q = _canonicalize_batch_stride(q) + kv = _canonicalize_batch_stride(kv) + ctx.scaling = scaling + valid_heads = _valid_head_mask(indices, q.size(-2)) + tl_out, tl_lse = sparse_mla_fwd_interface(q, kv, indices, sm_scale=scaling) + tl_out = _zero_invalid_heads(tl_out, valid_heads) + lse_zero = torch.zeros((), dtype=tl_lse.dtype, device=tl_lse.device) + tl_lse = torch.where(valid_heads, tl_lse, lse_zero) + + # Do not save tl_out/tl_lse: backward recomputes them just long enough to form + # delta and run the kernel. Saved inputs still go through autograd's saved-tensor + # hooks/offload path and retain_graph can recompute these tensors again. + ctx.save_for_backward(q, kv, indices, valid_heads) + + return tl_out, tl_lse + + @staticmethod + def backward(ctx, grad_output, grad_lse): + """ + Args: + grad_output: Gradient of the loss with respect to output + + Returns: + Gradients for q, kv, and indices (None for indices) + """ + q, kv, indices, valid_heads = ctx.saved_tensors + scaling = ctx.scaling + grad_output = grad_output.contiguous() + grad_output = _zero_invalid_heads(grad_output, valid_heads) + with torch.no_grad(): + tl_out, tl_lse = sparse_mla_fwd_interface(q, kv, indices, sm_scale=scaling) + tl_out = _zero_invalid_heads(tl_out, valid_heads) + lse_zero = torch.zeros((), dtype=tl_lse.dtype, device=tl_lse.device) + tl_lse = torch.where(valid_heads, tl_lse, lse_zero) + delta = sparse_mla_delta(tl_out, grad_output) + del tl_out + + tl_dq, tl_dkv = sparse_mla_bwd( + q, kv, None, grad_output, indices, tl_lse, sm_scale=scaling, delta=delta + ) + tl_dq = _zero_invalid_heads(tl_dq, valid_heads) + del tl_lse + + # Return gradients for each input (None for indices as it's not differentiable) + return tl_dq, tl_dkv, None, None + + +if not HAVE_TILELANG_SPARSE_MLA: + SparseMLA = None diff --git a/megatron/core/transformer/experimental_attention_variant/ops/tilelang_dsa.py b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_dsa.py new file mode 100644 index 00000000000..bedcdb4ca1f --- /dev/null +++ b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_dsa.py @@ -0,0 +1,908 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""TileLang-backed DSA hook implementations. + +This module keeps TileLang-specific batching, chunking, and sparse-KL streaming +out of the backend-neutral DSA control flow in ``dsa.py``. +""" + +from collections import OrderedDict +from typing import TYPE_CHECKING, Optional, Tuple + +import torch + +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.transformer.experimental_attention_variant import ( + dsa_indexer_loss, + dsa_layout, + dsa_masking, +) +from megatron.core.utils import get_pg_size + +if TYPE_CHECKING: + from megatron.core.packed_seq_params import PackedSeqParams + +try: + from megatron.core.transformer.experimental_attention_variant.ops.indexer import ( + lighting_indexer, + lighting_indexer_indices, + ) + from megatron.core.transformer.experimental_attention_variant.ops.tilelang_indexer_bwd import ( + is_supported_indexer_bwd_head_count, + ) +except (ImportError, OSError): + is_supported_indexer_bwd_head_count = None + lighting_indexer = None + lighting_indexer_indices = None + +try: + from megatron.core.transformer.experimental_attention_variant.ops.sparse_mla import SparseMLA +except (ImportError, OSError): + SparseMLA = None + +try: + from megatron.core.transformer.experimental_attention_variant.ops.tilelang_indexer_loss import ( + SparseIndexerKLLoss, + sparse_indexer_target_interface, + ) +except (ImportError, OSError): + SparseIndexerKLLoss = None + sparse_indexer_target_interface = None + + +# Reusable no-grad scratch buffers keyed by (name, shape, dtype, device). +_DSA_SCRATCH_CACHE_MAX_ENTRIES = 128 +_DSA_SCRATCH_CACHE_MAX_BYTES = 512 * 1024 * 1024 +_DSA_SCRATCH_CACHE = OrderedDict() +_DSA_SCRATCH_CACHE_TOTAL_BYTES = 0 + + +def _scratch_buffer_bytes(buf: torch.Tensor) -> int: + return buf.numel() * buf.element_size() + + +def _is_supported_sparse_mla_head_count(heads: int, kv_group: int = 1) -> bool: + """Return whether TileLang SparseMLA supports this query/KV head grouping. + + The forward and backward kernels pad ``head_kv`` to ``max(next_power_of_2(head_kv), 16)`` + and index the unpadded head dimension by that padded count with no head-dim bound, so they + only stay in bounds when no padding occurs. That requires ``head_kv`` (= ``heads // + kv_group``) to be a power of two and at least 16; any other value (e.g. 48, 192, or < 16) + must fall back to the unfused path rather than read/write past the real head count. + """ + if kv_group <= 0 or heads % kv_group != 0: + return False + head_kv = heads // kv_group + return head_kv >= 16 and (head_kv & (head_kv - 1)) == 0 + + +def _all_bfloat16(*tensors: torch.Tensor) -> bool: + return all(tensor.dtype == torch.bfloat16 for tensor in tensors) + + +def _evict_scratch_cache_if_needed() -> None: + """Bound scratch cache growth by LRU eviction.""" + global _DSA_SCRATCH_CACHE_TOTAL_BYTES + while ( + len(_DSA_SCRATCH_CACHE) > _DSA_SCRATCH_CACHE_MAX_ENTRIES + or _DSA_SCRATCH_CACHE_TOTAL_BYTES > _DSA_SCRATCH_CACHE_MAX_BYTES + ): + _, buf = _DSA_SCRATCH_CACHE.popitem(last=False) + _DSA_SCRATCH_CACHE_TOTAL_BYTES -= _scratch_buffer_bytes(buf) + + +def _get_scratch_buffer( + name: str, shape: Tuple[int, ...], dtype: torch.dtype, device: torch.device +) -> torch.Tensor: + """Get a reusable scratch tensor for temporary no-grad workspaces.""" + global _DSA_SCRATCH_CACHE_TOTAL_BYTES + key = (name, shape, dtype, device) + buf = _DSA_SCRATCH_CACHE.pop(key, None) + if buf is not None: + _DSA_SCRATCH_CACHE_TOTAL_BYTES -= _scratch_buffer_bytes(buf) + else: + buf = torch.empty(shape, dtype=dtype, device=device) + _DSA_SCRATCH_CACHE[key] = buf + _DSA_SCRATCH_CACHE_TOTAL_BYTES += _scratch_buffer_bytes(buf) + _evict_scratch_cache_if_needed() + return buf + + +def _topk_valid_mask( + topk_indices: torch.Tensor, starts: torch.Tensor, ends: torch.Tensor +) -> torch.Tensor: + """Compute the row-wise [start, end) validity mask for fused indexer outputs.""" + starts_for_cmp = starts.to(device=topk_indices.device, dtype=topk_indices.dtype).unsqueeze(-1) + ends_for_cmp = ends.to(device=topk_indices.device, dtype=topk_indices.dtype).unsqueeze(-1) + return (topk_indices >= starts_for_cmp) & (topk_indices < ends_for_cmp) + + +def _sanitize_fused_topk_indices( + topk_indices: torch.Tensor, starts: torch.Tensor, ends: torch.Tensor +) -> torch.Tensor: + """Mask fused indexer outputs in place and return the validity mask.""" + valid = _topk_valid_mask(topk_indices, starts, ends) + topk_indices.masked_fill_(~valid, -1) + return valid + + +def _sanitize_fused_topk_outputs( + topk_indices: torch.Tensor, + starts: torch.Tensor, + ends: torch.Tensor, + topk_scores: Optional[torch.Tensor] = None, +) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: + """Mask fused indexer outputs and optional scores to row-wise key bounds.""" + valid = _topk_valid_mask(topk_indices, starts, ends) + sanitized_indices = topk_indices.masked_fill(~valid, -1) + if topk_scores is not None: + topk_scores = topk_scores.masked_fill(~valid, float("-inf")) + return sanitized_indices, topk_scores + + +def _build_packed_cp_indexer_inputs( + index_k: torch.Tensor, + starts: torch.Tensor, + ends: torch.Tensor, + *, + packed_seq_params: "PackedSeqParams", + cp_size: int, + cp_rank: int, + single_packed_thd_sequence: bool, + local_query_start: int, + local_query_len: int, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Pack CP front/back key prefixes and translate query bounds to the packed key space.""" + if cp_size <= 1 or not 0 <= cp_rank < cp_size: + raise RuntimeError("packed CP TileLang indexer requires a valid CP rank and cp_size > 1") + if local_query_start < 0 or local_query_start + starts.numel() > local_query_len: + raise RuntimeError( + "packed CP TileLang indexer received an invalid local query slice: " + f"start={local_query_start}, rows={starts.numel()}, local_rows={local_query_len}" + ) + + cu_q, cu_k = dsa_layout.get_packed_qk_cu_seqlens(packed_seq_params) + device = index_k.device + sk = index_k.size(0) + layout = dsa_layout.build_packed_cp_indexer_layout( + cu_q.to(device=device), + cu_k.to(device=device), + cp_size=cp_size, + cp_rank=cp_rank, + key_size=sk, + local_key_layout=single_packed_thd_sequence and sk == local_query_len, + ) + segmented_k = index_k.index_select(0, layout.source_indices).contiguous() + + segment_ids_q = torch.repeat_interleave( + torch.arange(layout.segment_q_lengths.numel(), device=device), + layout.segment_q_lengths, + output_size=local_query_len, + ) + row_start = local_query_start + row_end = row_start + starts.numel() + row_segment_ids = segment_ids_q[row_start:row_end] + row_segment_starts = layout.segment_cu_k[:-1].index_select(0, row_segment_ids) + row_segment_ends = row_segment_starts + layout.segment_k_lengths.index_select( + 0, row_segment_ids + ) + row_global_starts = layout.segment_key_starts.index_select(0, row_segment_ids) + + local_starts = row_segment_starts + starts.to(torch.int64) - row_global_starts + local_ends = row_segment_starts + ends.to(torch.int64) - row_global_starts + local_starts = torch.maximum(local_starts, row_segment_starts) + local_starts = torch.minimum(local_starts, row_segment_ends) + local_ends = torch.maximum(local_ends, local_starts) + local_ends = torch.minimum(local_ends, row_segment_ends) + return ( + segmented_k, + local_starts.to(torch.int32).contiguous(), + local_ends.to(torch.int32).contiguous(), + layout.source_indices, + ) + + +def _remap_segmented_topk_indices( + topk_indices: torch.Tensor, source_indices: torch.Tensor +) -> torch.Tensor: + """Map valid indices from a segmented key tensor back to the original packed key tensor.""" + valid = topk_indices >= 0 + safe_indices = topk_indices.clamp(min=0).reshape(-1).to(torch.int64) + global_indices = source_indices.index_select(0, safe_indices).view_as(topk_indices) + return torch.where(valid, global_indices.to(topk_indices.dtype), topk_indices) + + +def fused_qk_topk_lighting( + q: torch.Tensor, + k: torch.Tensor, + weights: torch.Tensor, + index_topk: int, + starts: torch.Tensor, + ends: torch.Tensor, + block_size: int, + use_relu: bool = True, + use_local_indexer_varlen: bool = False, + single_packed_thd_sequence: bool = False, + local_packed_cp_rank: int = 0, + local_packed_cp_query_start: int = 0, + local_packed_cp_query_len: Optional[int] = None, + packed_seq_params: Optional["PackedSeqParams"] = None, + cp_size: int = 1, +) -> Optional[torch.Tensor]: + """Run fused TileLang indexer and return top-k indices [b, sq, topk].""" + if lighting_indexer_indices is None: + return None + if q.ndim != 4 or k.ndim != 3 or weights.ndim != 3: + return None + if not _all_bfloat16(q, k): + return None + + sq, b = q.size(0), q.size(1) + if k.size(1) != b or weights.size(1) != b: + return None + starts = starts.contiguous() + ends = ends.contiguous() + + topk_k = min(index_topk, k.size(0)) + topk_out = torch.empty((b, sq, topk_k), dtype=torch.int32, device=q.device) + for bi in range(b): + index_q = q[:, bi].contiguous() + index_k = k[:, bi].contiguous() + index_w = weights[:, bi].float().contiguous() + local_starts = starts + local_ends = ends + source_indices = None + if b == 1 and use_local_indexer_varlen and packed_seq_params is not None and cp_size > 1: + local_query_len = ( + local_packed_cp_query_len if local_packed_cp_query_len is not None else sq + ) + index_k, local_starts, local_ends, source_indices = _build_packed_cp_indexer_inputs( + index_k, + starts, + ends, + packed_seq_params=packed_seq_params, + cp_size=cp_size, + cp_rank=local_packed_cp_rank, + single_packed_thd_sequence=single_packed_thd_sequence, + local_query_start=local_packed_cp_query_start, + local_query_len=local_query_len, + ) + for start in range(0, sq, block_size): + end = min(start + block_size, sq) + topk_indices = lighting_indexer_indices( + index_q[start:end], + index_k, + index_w[start:end], + local_starts[start:end], + local_ends[start:end], + topk_k, + use_relu=use_relu, + ) + _sanitize_fused_topk_indices( + topk_indices, starts=local_starts[start:end], ends=local_ends[start:end] + ) + if source_indices is not None: + topk_indices = _remap_segmented_topk_indices(topk_indices, source_indices) + topk_out[bi, start:end].copy_(topk_indices) + + return topk_out + + +@torch.no_grad() +def _compute_topk_target_chunk_sum( + *, + query_h: torch.Tensor, + key_shared: Optional[torch.Tensor], + key_per_head: Optional[torch.Tensor], + s0: int, + s1: int, + idx_seq: torch.Tensor, + valid_seq: torch.Tensor, + softmax_scale: float, + head_chunk_size: int, + topk_chunk_size: int, + sk: int, + hn: int, +) -> torch.Tensor: + """Compute unnormalized target probability mass on top-k support for one sequence chunk.""" + s_len = s1 - s0 + topk = idx_seq.size(-1) + device = query_h.device + np = query_h.size(0) + + attn_chunk_sum = _get_scratch_buffer("kl_attn_chunk_sum", (s_len, topk), torch.float32, device) + attn_chunk_sum.zero_() + + for h0 in range(0, np, head_chunk_size): + h1 = min(h0 + head_chunk_size, np) + h_chunk = h1 - h0 + q_chunk = query_h[h0:h1, s0:s1, :] + q_chunk_float = q_chunk.float() + + if key_shared is None: + key_chunk = key_per_head[h0:h1] + flat_keys = key_chunk.reshape(h_chunk * sk, hn) + head_offsets = ( + torch.arange(h_chunk, device=device, dtype=torch.int64).view(-1, 1, 1) * sk + ) + else: + flat_keys = None + head_offsets = None + + # Two-pass online softmax over top-k chunks: + # 1) compute row-wise max and denominator; 2) recompute and accumulate probabilities. + # These accumulators are rebound to fresh tensors each top-k chunk, so they + # cannot reuse a scratch buffer in place. + running_max = torch.full( + (h_chunk, s_len), float("-inf"), dtype=torch.float32, device=device + ) + running_denom = torch.zeros((h_chunk, s_len), dtype=torch.float32, device=device) + + def _chunk_logits(idx_topk, valid_topk_chunk, k_len): + if key_shared is not None: + key_sel = key_shared.index_select(0, idx_topk.reshape(-1)).view(s_len, k_len, hn) + logits = torch.einsum("hsd,skd->hsk", q_chunk_float, key_sel.float()) + else: + flat_idx = idx_topk.unsqueeze(0) + head_offsets + key_sel = flat_keys.index_select(0, flat_idx.reshape(-1)).view( + h_chunk, s_len, k_len, hn + ) + logits = (q_chunk_float.unsqueeze(2) * key_sel.float()).sum(dim=-1) + logits = logits * softmax_scale + return logits.masked_fill(~valid_topk_chunk.unsqueeze(0), float("-inf")) + + for t0 in range(0, topk, topk_chunk_size): + t1 = min(t0 + topk_chunk_size, topk) + logits = _chunk_logits(idx_seq[:, t0:t1], valid_seq[:, t0:t1], t1 - t0) + chunk_max = logits.max(dim=-1).values + new_running_max = torch.maximum(running_max, chunk_max) + max_for_exp = torch.where( + torch.isfinite(new_running_max), new_running_max, torch.zeros_like(new_running_max) + ) + alpha = torch.exp(running_max - max_for_exp) + p_chunk = torch.exp(logits - max_for_exp.unsqueeze(-1)) + running_denom = running_denom * alpha + p_chunk.sum(dim=-1) + running_max = new_running_max + + stable_max = torch.where( + torch.isfinite(running_max), running_max, torch.zeros_like(running_max) + ) + inverse_denom = running_denom.clamp_min(1e-10).reciprocal() + for t0 in range(0, topk, topk_chunk_size): + t1 = min(t0 + topk_chunk_size, topk) + logits = _chunk_logits(idx_seq[:, t0:t1], valid_seq[:, t0:t1], t1 - t0) + probs = torch.exp(logits - stable_max.unsqueeze(-1)) * inverse_denom.unsqueeze(-1) + attn_chunk_sum[:, t0:t1] += probs.sum(dim=0) + + return attn_chunk_sum + + +def _compute_sparse_topk_kl_chunk( + target_chunk: torch.Tensor, index_logits_chunk: torch.Tensor, valid_seq: torch.Tensor +) -> torch.Tensor: + """Compute KL(target || index) sum for one [s_chunk, topk] chunk.""" + index_logits_chunk = index_logits_chunk.to(dtype=torch.float32, device=target_chunk.device) + target_chunk = target_chunk.to(dtype=torch.float32, device=index_logits_chunk.device) + with torch.no_grad(): + index_log_scores_chunk = dsa_masking.masked_log_softmax( + index_logits_chunk.detach(), valid_seq, dim=-1 + ) + index_scores_chunk = index_log_scores_chunk.exp().masked_fill(~valid_seq, 0.0) + kl_value = dsa_indexer_loss.indexer_kl_sum(target_chunk, index_log_scores_chunk, valid_seq) + grad_logits = (index_scores_chunk - target_chunk).masked_fill(~valid_seq, 0.0) + index_logits_for_grad = index_logits_chunk.masked_fill(~valid_seq, 0.0) + grad_surrogate = (index_logits_for_grad * grad_logits).sum() + return grad_surrogate + (kl_value - grad_surrogate).detach() + + +def _can_use_fused_sparse_indexer_target( + query: torch.Tensor, key: Optional[torch.Tensor], topk_indices: torch.Tensor +) -> bool: + """Return whether the fused TileLang target kernel supports these tensors.""" + return ( + sparse_indexer_target_interface is not None + and key is not None + and query.is_cuda + and key.is_cuda + and topk_indices.is_cuda + and query.ndim == 3 + and key.ndim == 2 + and query.dtype == torch.bfloat16 + and key.dtype == torch.bfloat16 + and query.size(-1) == key.size(-1) + and query.size(-1) % 16 == 0 + and topk_indices.ndim == 2 + and topk_indices.size(-1) % 64 == 0 + ) + + +def _can_use_fused_sparse_indexer_kl( + target: torch.Tensor, index_logits: torch.Tensor, valid_mask: torch.Tensor +) -> bool: + """Return whether the fused TileLang KL/score-gradient kernel supports these tensors.""" + return ( + SparseIndexerKLLoss is not None + and target.is_cuda + and index_logits.is_cuda + and valid_mask.is_cuda + and target.dtype == torch.float32 + and index_logits.dtype == torch.float32 + and valid_mask.dtype == torch.bool + and target.shape == index_logits.shape == valid_mask.shape + and target.ndim == 2 + and target.size(-1) % 256 == 0 + ) + + +def _canonicalize_topk_scores_for_tp_reduce( + topk_indices: torch.Tensor, topk_scores: torch.Tensor, *, sk: int +) -> Tuple[torch.Tensor, torch.Tensor]: + """Sort selected top-k slots by key index before slot-wise TP reductions.""" + valid = topk_indices >= 0 + topk_indices, topk_scores = dsa_masking.sort_topk_by_index( + topk_indices, valid, sk=sk, topk_scores=topk_scores + ) + assert topk_scores is not None + return topk_indices, topk_scores + + +def _accumulate_topk_kl_chunk( + *, + target_chunk: torch.Tensor, + index_logits_chunk: torch.Tensor, + valid_seq: torch.Tensor, + kl_sum: torch.Tensor, +) -> torch.Tensor: + """Normalize one target chunk and accumulate its sparse KL contribution.""" + if _can_use_fused_sparse_indexer_kl(target_chunk, index_logits_chunk, valid_seq): + return kl_sum + SparseIndexerKLLoss.apply( + target_chunk.contiguous(), index_logits_chunk.contiguous(), valid_seq.contiguous() + ) + normalized_target = dsa_indexer_loss.normalize_indexer_target_(target_chunk) + return kl_sum + _compute_sparse_topk_kl_chunk( + target_chunk=normalized_target, index_logits_chunk=index_logits_chunk, valid_seq=valid_seq + ) + + +def _stage_topk_target_chunk( + target_chunk: torch.Tensor, + *, + slot_prefix: str, + slot: int, + device: torch.device, + tp_group: torch.distributed.ProcessGroup, + tp_size: int, +) -> Tuple[torch.Tensor, Optional[torch.distributed.Work]]: + """Copy chunk into scratch slot and optionally launch async TP all-reduce.""" + target_chunk_work = _get_scratch_buffer( + f"{slot_prefix}_slot{slot}", tuple(target_chunk.shape), torch.float32, device + ) + target_chunk_work.copy_(target_chunk) + if tp_size > 1: + handle = torch.distributed.all_reduce(target_chunk_work, group=tp_group, async_op=True) + else: + handle = None + return target_chunk_work, handle + + +def _consume_pending_topk_kl_chunk( + *, + pending_handle: Optional[torch.distributed.Work], + pending_target_chunk: Optional[torch.Tensor], + pending_index_logits: Optional[torch.Tensor], + pending_valid_seq: Optional[torch.Tensor], + kl_sum: torch.Tensor, +) -> torch.Tensor: + """Finalize one pending chunk and accumulate its KL contribution into ``kl_sum``.""" + if pending_target_chunk is None: + return kl_sum + if pending_handle is not None: + pending_handle.wait() + return _accumulate_topk_kl_chunk( + target_chunk=pending_target_chunk, + index_logits_chunk=pending_index_logits, + valid_seq=pending_valid_seq, + kl_sum=kl_sum, + ) + + +def _enqueue_topk_kl_chunk( + *, + target_chunk: torch.Tensor, + index_logits_chunk: torch.Tensor, + valid_seq: torch.Tensor, + slot_prefix: str, + chunk_id: int, + device: torch.device, + tp_group: torch.distributed.ProcessGroup, + tp_size: int, + pending_handle: Optional[torch.distributed.Work], + pending_target_chunk: Optional[torch.Tensor], + pending_index_logits: Optional[torch.Tensor], + pending_valid_seq: Optional[torch.Tensor], + kl_sum: torch.Tensor, +) -> Tuple[ + torch.Tensor, + int, + Optional[torch.distributed.Work], + Optional[torch.Tensor], + Optional[torch.Tensor], + Optional[torch.Tensor], +]: + """Stage a new KL chunk, consume previous pending chunk, and update pending state.""" + slot = chunk_id & 1 + target_chunk_work, current_handle = _stage_topk_target_chunk( + target_chunk, + slot_prefix=slot_prefix, + slot=slot, + device=device, + tp_group=tp_group, + tp_size=tp_size, + ) + kl_sum = _consume_pending_topk_kl_chunk( + pending_handle=pending_handle, + pending_target_chunk=pending_target_chunk, + pending_index_logits=pending_index_logits, + pending_valid_seq=pending_valid_seq, + kl_sum=kl_sum, + ) + return (kl_sum, chunk_id + 1, current_handle, target_chunk_work, index_logits_chunk, valid_seq) + + +def fused_qk_topk_lighting_with_streaming_sparse_kl( + q: torch.Tensor, + k: torch.Tensor, + weights: torch.Tensor, + index_topk: int, + starts: torch.Tensor, + ends: torch.Tensor, + block_size: int, + query: torch.Tensor, + key: torch.Tensor, + softmax_scale: float, + loss_coeff: float, + pg_collection: ProcessGroupCollection, + query_valid_rows: Optional[torch.Tensor] = None, + calculate_per_token_loss: bool = False, + seq_chunk_size: int = 512, + head_chunk_size: int = 16, + topk_chunk_size: int = 1024, + use_relu: bool = True, + use_local_indexer_varlen: bool = False, + single_packed_thd_sequence: bool = False, + local_packed_cp_rank: int = 0, + local_packed_cp_query_start: int = 0, + local_packed_cp_query_len: Optional[int] = None, + packed_seq_params: Optional["PackedSeqParams"] = None, + cp_size: int = 1, +) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: + """Run the fused TileLang indexer with streaming sparse KL accumulation. + + The objective matches ``compute_dsa_indexer_loss`` on the selected top-k support. TileLang + streams query/head/top-k chunks and overlaps TP target reduction to avoid materializing dense + scores; its custom gradient surrogate supplies the same log-softmax gradient for fused indexer + logits. Target normalization, KL evaluation, and token reduction use the shared backend-neutral + helpers in ``dsa_indexer_loss``. + """ + if lighting_indexer is None: + return None + if q.ndim != 4 or k.ndim != 3 or weights.ndim != 3: + return None + if not _all_bfloat16(q, k): + return None + if is_supported_indexer_bwd_head_count is None or not is_supported_indexer_bwd_head_count( + q.size(2) + ): + return None + + query, _ = dsa_layout.ensure_sbhd(query, "query") + key, _ = dsa_layout.ensure_sbhd(key, "key") + sq, b = q.size(0), q.size(1) + sq_q, b_q, np, hn = query.size() + sk, b_k, nk, hk = key.size() + if k.size(1) != b or weights.size(1) != b: + return None + if sq_q != sq or b_q != b or b_k != b or hk != hn: + return None + if nk != 1 and nk != np: + return None + query_valid_rows = dsa_masking.normalize_query_valid_rows( + query_valid_rows, b=b, sq=sq, device=query.device + ) + + starts = starts.contiguous() + ends = ends.contiguous() + + topk_out = None + kl_sum = torch.zeros((), dtype=torch.float32, device=q.device) + tp_size = get_pg_size(pg_collection.tp) + pending_handle = None + pending_target_chunk = None + pending_index_logits = None + pending_valid_seq = None + chunk_id = 0 + for bi in range(b): + query_h = query[:, bi].permute(1, 0, 2).contiguous() + if nk == 1: + key_shared = key[:, bi, 0].contiguous() + key_per_head = None + else: + key_shared = None + key_per_head = key[:, bi].permute(1, 0, 2).contiguous() + + index_q = q[:, bi].contiguous() + index_k = k[:, bi].contiguous() + index_w = weights[:, bi].float().contiguous() + local_starts = starts + local_ends = ends + source_indices = None + if b == 1 and use_local_indexer_varlen and packed_seq_params is not None and cp_size > 1: + local_query_len = ( + local_packed_cp_query_len if local_packed_cp_query_len is not None else sq + ) + index_k, local_starts, local_ends, source_indices = _build_packed_cp_indexer_inputs( + index_k, + starts, + ends, + packed_seq_params=packed_seq_params, + cp_size=cp_size, + cp_rank=local_packed_cp_rank, + single_packed_thd_sequence=single_packed_thd_sequence, + local_query_start=local_packed_cp_query_start, + local_query_len=local_query_len, + ) + + for start in range(0, sq, block_size): + end = min(start + block_size, sq) + topk_scores, topk_indices = lighting_indexer( + index_q[start:end], + index_k, + index_w[start:end], + local_starts[start:end], + local_ends[start:end], + min(index_topk, k.size(0)), + topk_indices=None, + use_relu=use_relu, + ) + topk_indices, topk_scores = _sanitize_fused_topk_outputs( + topk_indices=topk_indices, + starts=local_starts[start:end], + ends=local_ends[start:end], + topk_scores=topk_scores, + ) + if source_indices is not None: + topk_indices = _remap_segmented_topk_indices(topk_indices, source_indices) + if tp_size > 1: + topk_indices, topk_scores = _canonicalize_topk_scores_for_tp_reduce( + topk_indices, topk_scores, sk=sk + ) + + if topk_out is None: + topk_out = torch.empty( + (b, sq, topk_indices.size(-1)), + dtype=topk_indices.dtype, + device=topk_indices.device, + ) + topk_out[bi, start:end].copy_(topk_indices) + + s_len = end - start + for rel_start in range(0, s_len, seq_chunk_size): + rel_end = min(rel_start + seq_chunk_size, s_len) + abs_start = start + rel_start + abs_end = start + rel_end + + idx_seq_raw = topk_indices[rel_start:rel_end].to(device=query.device) + valid_seq = idx_seq_raw >= 0 + if query_valid_rows is not None: + row_valid = query_valid_rows[bi, abs_start:abs_end] + valid_seq = valid_seq & row_valid.unsqueeze(-1) + loss_topk_indices = idx_seq_raw.masked_fill(~valid_seq, -1).contiguous() + query_chunk = query[abs_start:abs_end, bi].contiguous() + if _can_use_fused_sparse_indexer_target(query_chunk, key_shared, loss_topk_indices): + target_chunk = sparse_indexer_target_interface( + query_chunk, key_shared, loss_topk_indices, softmax_scale + ) + else: + target_chunk = _compute_topk_target_chunk_sum( + query_h=query_h, + key_shared=key_shared, + key_per_head=key_per_head, + s0=abs_start, + s1=abs_end, + idx_seq=idx_seq_raw.clamp(min=0).to(torch.int64), + valid_seq=valid_seq, + softmax_scale=softmax_scale, + head_chunk_size=head_chunk_size, + topk_chunk_size=topk_chunk_size, + sk=sk, + hn=hn, + ) + index_logits_chunk = topk_scores[rel_start:rel_end] + ( + kl_sum, + chunk_id, + pending_handle, + pending_target_chunk, + pending_index_logits, + pending_valid_seq, + ) = _enqueue_topk_kl_chunk( + target_chunk=target_chunk, + index_logits_chunk=index_logits_chunk, + valid_seq=valid_seq, + slot_prefix="stream_kl_target", + chunk_id=chunk_id, + device=query.device, + tp_group=pg_collection.tp, + tp_size=tp_size, + pending_handle=pending_handle, + pending_target_chunk=pending_target_chunk, + pending_index_logits=pending_index_logits, + pending_valid_seq=pending_valid_seq, + kl_sum=kl_sum, + ) + kl_sum = _consume_pending_topk_kl_chunk( + pending_handle=pending_handle, + pending_target_chunk=pending_target_chunk, + pending_index_logits=pending_index_logits, + pending_valid_seq=pending_valid_seq, + kl_sum=kl_sum, + ) + + if topk_out is None: + return None + valid_row_count = query_valid_rows.sum() if query_valid_rows is not None else None + kl_div = dsa_indexer_loss.reduce_indexer_kl_sum( + kl_sum, + num_rows=b * sq, + calculate_per_token_loss=calculate_per_token_loss, + valid_row_count=valid_row_count, + ) + return topk_out, kl_div * loss_coeff + + +def fused_sparse_mla_absorbed( + query: torch.Tensor, + key: torch.Tensor, + topk_indices: torch.Tensor, + softmax_scale: float, + v_channels: int, +) -> Optional[torch.Tensor]: + """Run fused SparseMLA kernel for absorbed-MLA path.""" + if SparseMLA is None: + return None + + if query.ndim != 4 or key.ndim != 4 or topk_indices.ndim != 3: + return None + if not _all_bfloat16(query, key): + return None + if key.size(2) != 1: + return None + if query.size(1) != key.size(1) or topk_indices.size(0) != query.size(1): + return None + if topk_indices.size(1) != query.size(0): + return None + if query.size(-1) != key.size(-1): + return None + if query.size(-1) != 576 or v_channels != 512: + # Current copied TileLang kernels are specialized for GLM5/DeepSeek V3.2 absorbed dims. + return None + query_heads = query.size(2) + if query_heads <= 0: + return None + kernel_heads = max(query_heads, 16) + if not _is_supported_sparse_mla_head_count(kernel_heads, kv_group=key.size(2)): + return None + if topk_indices.size(-1) % 64 != 0: + return None + + query_bshd = query.permute(1, 0, 2, 3).contiguous() + if kernel_heads != query_heads: + # SparseMLA uses a minimum 16-head tile without head bounds. Pad the caller + # tensor so small TP shards stay in bounds, then discard those heads below. + query_bshd = torch.nn.functional.pad(query_bshd, (0, 0, 0, kernel_heads - query_heads)) + key_bshd = key.permute(1, 0, 2, 3).contiguous() + indices_bsgk = topk_indices.unsqueeze(2).to(torch.int32).contiguous() + out, _ = SparseMLA.apply(query_bshd, key_bshd, indices_bsgk, softmax_scale) + if out.ndim != 4 or out.size(2) != kernel_heads or out.size(-1) != v_channels: + return None + out = out[:, :, :query_heads] + return out.permute(1, 0, 2, 3).contiguous() + + +def run_fused_qk_topk( + q: torch.Tensor, + k: torch.Tensor, + weights: torch.Tensor, + index_topk: int, + starts: torch.Tensor, + ends: torch.Tensor, + block_size: int, + use_relu: bool = True, + use_local_indexer_varlen: bool = False, + single_packed_thd_sequence: bool = False, + local_packed_cp_rank: int = 0, + local_packed_cp_query_start: int = 0, + local_packed_cp_query_len: Optional[int] = None, + packed_seq_params: Optional["PackedSeqParams"] = None, + cp_size: int = 1, +) -> Optional[torch.Tensor]: + """Optional fused indexer hook backed by TileLang.""" + return fused_qk_topk_lighting( + q, + k, + weights, + index_topk, + starts, + ends, + block_size, + use_relu, + use_local_indexer_varlen=use_local_indexer_varlen, + single_packed_thd_sequence=single_packed_thd_sequence, + local_packed_cp_rank=local_packed_cp_rank, + local_packed_cp_query_start=local_packed_cp_query_start, + local_packed_cp_query_len=local_packed_cp_query_len, + packed_seq_params=packed_seq_params, + cp_size=cp_size, + ) + + +def run_fused_qk_topk_with_loss( + q: torch.Tensor, + k: torch.Tensor, + weights: torch.Tensor, + index_topk: int, + starts: torch.Tensor, + ends: torch.Tensor, + block_size: int, + query: torch.Tensor, + key: torch.Tensor, + softmax_scale: float, + loss_coeff: float, + pg_collection: ProcessGroupCollection, + query_valid_rows: Optional[torch.Tensor] = None, + calculate_per_token_loss: bool = False, + use_relu: bool = True, + use_local_indexer_varlen: bool = False, + single_packed_thd_sequence: bool = False, + local_packed_cp_rank: int = 0, + local_packed_cp_query_start: int = 0, + local_packed_cp_query_len: Optional[int] = None, + packed_seq_params: Optional["PackedSeqParams"] = None, + cp_size: int = 1, +) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: + """Optional fused indexer+loss hook backed by TileLang.""" + return fused_qk_topk_lighting_with_streaming_sparse_kl( + q=q, + k=k, + weights=weights, + index_topk=index_topk, + starts=starts, + ends=ends, + block_size=block_size, + query=query, + key=key, + softmax_scale=softmax_scale, + loss_coeff=loss_coeff, + pg_collection=pg_collection, + query_valid_rows=query_valid_rows, + calculate_per_token_loss=calculate_per_token_loss, + use_relu=use_relu, + use_local_indexer_varlen=use_local_indexer_varlen, + single_packed_thd_sequence=single_packed_thd_sequence, + local_packed_cp_rank=local_packed_cp_rank, + local_packed_cp_query_start=local_packed_cp_query_start, + local_packed_cp_query_len=local_packed_cp_query_len, + packed_seq_params=packed_seq_params, + cp_size=cp_size, + ) + + +def run_fused_absorbed_sparse_attention( + query: torch.Tensor, + key: torch.Tensor, + topk_indices: torch.Tensor, + softmax_scale: float, + v_channels: int, +) -> Optional[torch.Tensor]: + """Optional fused sparse-attention hook backed by TileLang.""" + return fused_sparse_mla_absorbed(query, key, topk_indices, softmax_scale, v_channels) diff --git a/megatron/core/transformer/experimental_attention_variant/ops/tilelang_indexer_bwd.py b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_indexer_bwd.py new file mode 100644 index 00000000000..33c5627728b --- /dev/null +++ b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_indexer_bwd.py @@ -0,0 +1,234 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# ruff: noqa +# Adapted from: +# https://github.com/tile-ai/tilelang/blob/4956b5835fa554af6c03d4a6289cad44bf310869/ +# examples/dsa_sparse_finetune/indexer_bwd.py +import threading +from collections import OrderedDict + +import torch + +from .tilelang_utils import ( + HAVE_TILELANG, + T, + _get_cached_kernel, + _next_power_of_two, + _round_up, + require_tilelang, +) +from .tilelang_utils import tilelang as tl +from .tilelang_utils import tilelang_jit + +BF16 = T.bfloat16 if HAVE_TILELANG else None +FP32 = T.float32 if HAVE_TILELANG else None +INT32 = T.int32 if HAVE_TILELANG else None +_tilelang_indexer_bwd_kernel_cache = OrderedDict() +_tilelang_indexer_bwd_cache_lock = threading.Lock() + +if HAVE_TILELANG: + pass_configs = { + tl.PassConfigKey.TL_DISABLE_TMA_LOWER: True, + tl.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, + } +else: + pass_configs = {} + + +def _canonical_topk(topk: int, block_i: int = 32) -> int: + return _round_up(_next_power_of_two(topk), block_i) + + +def is_supported_indexer_bwd_head_count(heads: int) -> bool: + """Return whether TileLang indexer backward supports this indexer head count.""" + return heads <= 64 and heads % 8 == 0 + + +def _get_indexer_bwd_kernel(heads: int, dim: int, topk: int, use_relu: bool = True): + num_threads = 32 if heads < 16 else 128 + return _get_cached_kernel( + _tilelang_indexer_bwd_kernel_cache, + _tilelang_indexer_bwd_cache_lock, + (heads, dim, topk, use_relu), + lambda: tl_indexer_bwd_impl(heads, dim, topk, num_threads=num_threads, use_relu=use_relu), + ) + + +@tilelang_jit(pass_configs=pass_configs) +def tl_indexer_bwd_impl( # pragma: no cover + heads: int, + dim: int, + topk: int, + block_I: int = 32, + num_stages: int = 0, + num_threads: int = 128, + use_relu: bool = True, +): + """Build tilelang backward kernel for sparse indexer.""" + require_tilelang() + assert num_stages == 0 + assert topk == tl.math.next_power_of_2(topk) + assert topk % block_I == 0 + assert heads <= 64 and heads % 8 == 0 + seq_len = T.symbolic("seq_len") + q_seq_len = T.symbolic("q_seq_len") + + dtype: str = BF16 + accum_dtype: str = FP32 + index_q_shape = [q_seq_len, heads, dim] + weights_shape = [q_seq_len, heads] + index_k_shape = [seq_len, dim] + shape_p = [q_seq_len, topk] + topk_indices_shape = [q_seq_len, topk] + + pad_heads = heads + if heads < 16: + pad_heads = 16 + + @T.prim_func + def tl_indexer_bwd_kernel( + IndexQ: T.Tensor(index_q_shape, dtype), + IndexK: T.Tensor(index_k_shape, dtype), + Weights: T.Tensor(weights_shape, FP32), + TopkIndices: T.Tensor(topk_indices_shape, INT32), + OGrad: T.Tensor(shape_p, FP32), + dIndexQ: T.Tensor(index_q_shape, dtype), + dWeights: T.Tensor(weights_shape, FP32), + dIndexK: T.Tensor(index_k_shape, FP32), + ): + + with T.Kernel(q_seq_len, threads=num_threads) as (bx): + index_q_shared = T.alloc_shared([pad_heads, dim], dtype=FP32) + weights_shared = T.alloc_shared([pad_heads], dtype=FP32) + index_k_shared = T.alloc_shared([block_I, dim], dtype=FP32) + indices_shared = T.alloc_shared([block_I], dtype=INT32) + d_index_q_frag = T.alloc_fragment([pad_heads, dim], dtype=accum_dtype) + d_weights_frag = T.alloc_fragment([pad_heads], dtype=accum_dtype) + d_index_k_frag = T.alloc_fragment([block_I, dim], dtype=accum_dtype) + logits = T.alloc_fragment((block_I, pad_heads), dtype=accum_dtype) + _logits = T.alloc_shared((block_I, pad_heads), dtype=accum_dtype) + grad = T.alloc_shared([block_I], dtype=FP32) + + num_blocks = T.ceildiv(topk, block_I) + for i, j in T.Parallel(pad_heads, dim): + index_q_shared[i, j] = T.if_then_else(i < heads, IndexQ[bx, i, j], 0) + for i in T.Parallel(heads): + weights_shared[i] = Weights[bx, i] + + T.fill(d_index_q_frag, 0) + T.fill(d_weights_frag, 0) + + for bi_i in T.serial(num_blocks): + for i in T.Parallel(block_I): + if bi_i * block_I + i < topk: + indices_shared[i] = TopkIndices[bx, bi_i * block_I + i] + grad[i] = OGrad[bx, bi_i * block_I + i] + + T.sync_threads() + for i, j in T.Parallel(block_I, dim): + index_k_shared[i, j] = T.if_then_else( + indices_shared[i] > -1 and indices_shared[i] < seq_len, + IndexK[indices_shared[i], j], + 0, + ) + + T.sync_threads() + T.gemm( + index_k_shared, + index_q_shared, + logits, + transpose_A=False, + transpose_B=True, + clear_accum=True, + ) + d_weights_i = T.alloc_fragment((block_I, pad_heads), accum_dtype) + for i, j in T.Parallel(block_I, heads): + d_weights_i[i, j] = grad[i] * ( + T.max(logits[i, j], 0) if use_relu else logits[i, j] + ) + T.reduce_sum(d_weights_i, d_weights_frag, dim=0, clear=False) + + for i, j in T.Parallel(block_I, pad_heads): + _logits[i, j] = T.if_then_else( + (logits[i, j] > 0 if use_relu else True) and j < heads, + grad[i] * weights_shared[j], + 0, + ) + T.sync_threads() + T.gemm( + _logits, + index_k_shared, + d_index_q_frag, + transpose_A=True, + transpose_B=False, + clear_accum=False, + ) + + T.gemm( + _logits, + index_q_shared, + d_index_k_frag, + transpose_A=False, + transpose_B=False, + clear_accum=True, + ) + + for i, j in T.Parallel(block_I, dim): + if indices_shared[i] > -1 and indices_shared[i] < seq_len: + T.atomic_add(dIndexK[indices_shared[i], j], d_index_k_frag[i, j]) + + T.copy(d_index_q_frag[:heads, :], dIndexQ[bx, :, :]) + T.copy(d_weights_frag[:heads], dWeights[bx, :]) + + return tl_indexer_bwd_kernel + + +def indexer_bwd_interface( # pragma: no cover + index_q: torch.Tensor, + weights: torch.Tensor, + index_k: torch.Tensor, + topk_indices: torch.Tensor, + grad_scores: torch.Tensor, + use_relu: bool = True, +): + """Run indexer backward kernel and return gradients for q/w/k.""" + require_tilelang() + _, head_num, head_dim = index_q.shape + k_top = int(topk_indices.shape[1]) + assert k_top > 0, "topk must be positive" + padded_topk = _canonical_topk(k_top) + + if padded_topk != k_top: + padded_indices = torch.full( + (topk_indices.size(0), padded_topk), + -1, + dtype=topk_indices.dtype, + device=topk_indices.device, + ) + padded_indices[:, :k_top].copy_(topk_indices) + topk_indices = padded_indices + + padded_grad_scores = torch.zeros( + (grad_scores.size(0), padded_topk), dtype=grad_scores.dtype, device=grad_scores.device + ) + padded_grad_scores[:, :k_top].copy_(grad_scores) + grad_scores = padded_grad_scores + + grad_scores = grad_scores.contiguous() + weights_kernel = weights.to(dtype=torch.float32).contiguous() + grad_q = torch.empty_like(index_q) + grad_w = torch.empty_like(weights, dtype=torch.float32) + grad_k = torch.zeros_like(index_k, dtype=torch.float32) + + bwd_kernel = _get_indexer_bwd_kernel(head_num, head_dim, padded_topk, use_relu=use_relu) + bwd_kernel( + index_q.contiguous(), + index_k.contiguous(), + weights_kernel, + topk_indices.contiguous(), + grad_scores, + grad_q, + grad_w, + grad_k, + ) + + return grad_q, grad_w, grad_k.to(index_k.dtype) diff --git a/megatron/core/transformer/experimental_attention_variant/ops/tilelang_indexer_fwd.py b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_indexer_fwd.py new file mode 100644 index 00000000000..ffa45ccdf90 --- /dev/null +++ b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_indexer_fwd.py @@ -0,0 +1,224 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# ruff: noqa +# Adapted from: +# https://github.com/tile-ai/tilelang/blob/4956b5835fa554af6c03d4a6289cad44bf310869/ +# examples/deepseek_v32/fp8_lighting_indexer.py +import threading +from collections import OrderedDict + +import torch + +from .tilelang_utils import ( + HAVE_TILELANG, + T, + _get_cached_kernel, + require_tilelang, + tilelang, + tilelang_jit, +) + +_tilelang_indexer_fwd_kernel_cache = OrderedDict() +_tilelang_indexer_clean_logits_kernel_cache = OrderedDict() +_tilelang_indexer_fwd_cache_lock = threading.Lock() + + +def _get_clean_logits_kernel(threads: int = 512, block_K: int = 4096): + return _get_cached_kernel( + _tilelang_indexer_clean_logits_kernel_cache, + _tilelang_indexer_fwd_cache_lock, + (threads, block_K), + lambda: clean_logits_(threads=threads, block_K=block_K), + ) + + +def _get_indexer_fwd_kernel( + heads: int, + index_dim: int, + block_N: int = 256, + num_stages: int = 3, + threads: int = 512, + use_relu: bool = True, +): + return _get_cached_kernel( + _tilelang_indexer_fwd_kernel_cache, + _tilelang_indexer_fwd_cache_lock, + (heads, index_dim, block_N, num_stages, threads, use_relu), + lambda: tl_indexer_fwd_impl( + heads=heads, + index_dim=index_dim, + block_N=block_N, + num_stages=num_stages, + threads=threads, + use_relu=use_relu, + ), + ) + + +_TL_INDEXER_FWD_PASS_CONFIGS = ( + {tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True} if HAVE_TILELANG else {} +) + + +@tilelang_jit(pass_configs=_TL_INDEXER_FWD_PASS_CONFIGS) +def tl_indexer_fwd_impl( # pragma: no cover + heads, index_dim, block_N=256, num_stages=3, threads=512, block_Q=None, use_relu=True +): + """Build tilelang forward kernel for sparse indexer logits.""" + require_tilelang() + assert heads > 0 + if block_Q is None: + block_Q = max(1, 128 // heads) + dtype = T.bfloat16 + accum_dtype = T.float32 + index_dtype = T.int32 + + seq_len = T.dynamic("seq_len") + seq_len_kv = T.dynamic("seq_len_kv") + + index_q_shape = [seq_len * heads, index_dim] + index_k_shape = [seq_len_kv, index_dim] + logits_shape = [seq_len, seq_len_kv] + + @T.prim_func + def tl_indexer_fwd_kernel( + IndexQ: T.Tensor(index_q_shape, dtype), # type: ignore + IndexK: T.Tensor(index_k_shape, dtype), # type: ignore + Logits: T.Tensor(logits_shape, accum_dtype), # type: ignore + Weights: T.Tensor([seq_len, heads], accum_dtype), # type: ignore + CuSeqLenKS: T.Tensor([seq_len], index_dtype), # type: ignore + CuSeqLenKE: T.Tensor([seq_len], index_dtype), # type: ignore + ): + with T.Kernel(T.ceildiv(seq_len, block_Q), threads=threads) as bx: + index_q_shared = T.alloc_shared([block_Q * heads, index_dim], dtype) + index_k_shared = T.alloc_shared([block_N, index_dim], dtype) + s = T.alloc_fragment([block_N, block_Q * heads], accum_dtype) + s_reshaped = T.reshape(s, (block_N, block_Q, heads)) + logits_shared = T.alloc_shared([block_N, block_Q], accum_dtype) + weights = T.alloc_fragment([block_Q, heads], accum_dtype) + + seq_len_i = bx * block_Q + + cu_k_s_min = T.alloc_var(index_dtype) + cu_k_e_max = T.alloc_var(index_dtype) + + cu_k_s_min = 2147483647 + cu_k_e_max = -2147483648 + + for bq_i in T.serial(block_Q): + q_idx = seq_len_i + bq_i + if q_idx < seq_len: + k_s = T.max(T.min(CuSeqLenKS[q_idx], seq_len_kv), 0) + cu_k_s_min = T.min(cu_k_s_min, k_s) + for bq_i in T.serial(block_Q): + q_idx = seq_len_i + bq_i + if q_idx < seq_len: + k_e = T.max(T.min(CuSeqLenKE[q_idx], seq_len_kv), 0) + cu_k_e_max = T.max(cu_k_e_max, k_e) + + # Clamp bounds to [0, seq_len_kv] and normalize empty rows. + cu_k_s_min = T.max(cu_k_s_min, 0) + cu_k_s_min = T.min(cu_k_s_min, seq_len_kv) + cu_k_e_max = T.max(cu_k_e_max, 0) + cu_k_e_max = T.min(cu_k_e_max, seq_len_kv) + if cu_k_e_max < cu_k_s_min: + cu_k_e_max = cu_k_s_min + + for bq_i, h_i, d_i in T.Parallel(block_Q, heads, index_dim): + q_idx = seq_len_i + bq_i + index_q_shared[bq_i * heads + h_i, d_i] = T.if_then_else( + q_idx < seq_len, IndexQ[q_idx * heads + h_i, d_i], 0 + ) + for bq_i, h_i in T.Parallel(block_Q, heads): + q_idx = seq_len_i + bq_i + weights[bq_i, h_i] = T.if_then_else(q_idx < seq_len, Weights[q_idx, h_i], 0) + + for nbn_i in T.Pipelined( + T.ceildiv(cu_k_e_max - cu_k_s_min, block_N), num_stages=num_stages + ): + for bn_i, d_i in T.Parallel(block_N, index_dim): + k_idx = cu_k_s_min + nbn_i * block_N + bn_i + index_k_shared[bn_i, d_i] = T.if_then_else( + k_idx >= 0 and k_idx < cu_k_e_max, IndexK[k_idx, d_i], 0 + ) + + T.gemm( + index_k_shared, + index_q_shared, + s, + transpose_B=True, + clear_accum=True, + policy=T.GemmWarpPolicy.FullCol, + ) + + for bn_i, bq_i, h_i in T.Parallel(block_N, block_Q, heads): + s_reshaped[bn_i, bq_i, h_i] = ( + T.max(s_reshaped[bn_i, bq_i, h_i], 0) + if use_relu + else s_reshaped[bn_i, bq_i, h_i] + ) * weights[bq_i, h_i] + + T.reduce_sum(s_reshaped, logits_shared, dim=-1, clear=True) + + # Keep this write deterministic to satisfy data-race verification. + for bq_i in T.serial(block_Q): + q_idx = seq_len_i + bq_i + if q_idx < seq_len: + for bn_i in T.serial(block_N): + k_idx = cu_k_s_min + nbn_i * block_N + bn_i + if k_idx >= 0 and k_idx < cu_k_e_max: + Logits[q_idx, k_idx] = logits_shared[bn_i, bq_i] + + return tl_indexer_fwd_kernel + + +@tilelang_jit +def clean_logits_(threads: int = 512, block_K: int = 4096): # pragma: no cover + """Build kernel that masks out invalid key ranges in logits.""" + require_tilelang() + seq_len = T.dynamic("seq_len") + seq_len_kv = T.dynamic("seq_len_kv") + + dtype = T.float + indices_dtype = T.int32 + + @T.prim_func + def clean_logits_kernel( + Logits: T.Tensor([seq_len, seq_len_kv], dtype), # type: ignore + CuSeqLenKS: T.Tensor([seq_len], indices_dtype), # type: ignore + CuSeqLenKE: T.Tensor([seq_len], indices_dtype), # type: ignore + ): + with T.Kernel(seq_len, threads=threads) as bx: + tx = T.thread_binding(0, threads, thread="threadIdx.x") + cu_k_s = CuSeqLenKS[bx] + cu_k_e = CuSeqLenKE[bx] + + for n_i in T.Pipelined(T.ceildiv(seq_len_kv, block_K)): + for k_i in T.serial(block_K // threads): + idx = n_i * block_K + k_i * threads + tx + if idx < seq_len_kv and (idx < cu_k_s or idx >= cu_k_e): + Logits[bx, idx] = -T.infinity(dtype) + + return clean_logits_kernel + + +def indexer_fwd_interface( # pragma: no cover + q, kv, weights, cu_seqlen_ks, cu_seqlen_ke, clean_logits=True, use_relu=True +): + """Run indexer forward kernel and optionally clean logits by row bounds.""" + require_tilelang() + seq_len, heads, index_dim = q.shape + seq_len_kv = kv.shape[0] + weights = weights.to(dtype=torch.float32).contiguous() + + tl_indexer_fwd_kernel = _get_indexer_fwd_kernel( + heads=heads, index_dim=index_dim, use_relu=use_relu + ) + logits = torch.empty([seq_len, seq_len_kv], device=q.device, dtype=torch.float32) + tl_indexer_fwd_kernel( + q.view(seq_len * heads, index_dim), kv, logits, weights, cu_seqlen_ks, cu_seqlen_ke + ) + + if clean_logits: + clean_logits_kernel = _get_clean_logits_kernel() + clean_logits_kernel(logits, cu_seqlen_ks, cu_seqlen_ke) + return logits diff --git a/megatron/core/transformer/experimental_attention_variant/ops/tilelang_indexer_loss.py b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_indexer_loss.py new file mode 100644 index 00000000000..3bfbfd70c2a --- /dev/null +++ b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_indexer_loss.py @@ -0,0 +1,344 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""TileLang kernels for the sparse DSA indexer KL target and score gradient.""" + +import threading +from collections import OrderedDict + +import torch + +from .tilelang_utils import ( + HAVE_TILELANG, + T, + _get_cached_kernel, + _normalize_sm_scale, + require_tilelang, + tilelang, + tilelang_jit, +) + +_target_kernel_cache = OrderedDict() +_kl_kernel_cache = OrderedDict() +_kernel_cache_lock = threading.Lock() + +_PASS_CONFIGS = {tilelang.PassConfigKey.TL_ENABLE_FAST_MATH: True} if HAVE_TILELANG else {} + + +def _get_target_kernel( + heads: int, + dim: int, + topk: int, + softmax_scale: float, + block_h: int = 32, + block_i: int = 64, + num_stages: int = 2, + threads: int = 256, +): + scale = _normalize_sm_scale(softmax_scale) + key = (heads, dim, topk, scale, block_h, block_i, num_stages, threads) + return _get_cached_kernel( + _target_kernel_cache, + _kernel_cache_lock, + key, + lambda: sparse_indexer_target( + heads=heads, + dim=dim, + topk=topk, + softmax_scale=scale, + block_h=block_h, + block_i=block_i, + num_stages=num_stages, + threads=threads, + ), + ) + + +def _get_kl_kernel(topk: int, block_i: int = 256, threads: int = 256): + key = (topk, block_i, threads) + return _get_cached_kernel( + _kl_kernel_cache, + _kernel_cache_lock, + key, + lambda: sparse_indexer_kl(topk=topk, block_i=block_i, threads=threads), + ) + + +@tilelang_jit(out_idx=[-1], pass_configs=_PASS_CONFIGS) +def sparse_indexer_target( # pragma: no cover + heads: int, + dim: int, + topk: int, + softmax_scale: float, + block_h: int = 32, + block_i: int = 64, + num_stages: int = 2, + threads: int = 256, +): + """Build a kernel that sums selected-key attention probabilities over local heads.""" + require_tilelang() + assert heads > 0 + assert dim % 16 == 0 + assert topk % block_i == 0 + + seq_len = T.dynamic("seq_len") + seq_len_kv = T.dynamic("seq_len_kv") + dtype = T.bfloat16 + accum_dtype = T.float32 + index_dtype = T.int32 + num_tiles = tilelang.cdiv(topk, block_i) + num_head_tiles = tilelang.cdiv(heads, block_h) + scale_log2 = softmax_scale * 1.4426950408889634 + + @T.prim_func + def main( + Query: T.Tensor([seq_len, heads, dim], dtype), # type: ignore + Key: T.Tensor([seq_len_kv, dim], dtype), # type: ignore + Indices: T.Tensor([seq_len, topk], index_dtype), # type: ignore + Target: T.Tensor([seq_len, topk], accum_dtype), # type: ignore + ): + with T.Kernel(seq_len, threads=threads) as row: + query_shared = T.alloc_shared([block_h, dim], dtype) + key_shared = T.alloc_shared([block_i, dim], dtype) + scores = T.alloc_fragment([block_h, block_i], accum_dtype) + probabilities = T.alloc_fragment([block_h, block_i], accum_dtype) + target_tile = T.alloc_fragment([block_i], accum_dtype) + valid = T.alloc_fragment([block_i], "bool") + row_max = T.alloc_fragment([block_h], accum_dtype) + previous_max = T.alloc_fragment([block_h], accum_dtype) + tile_max = T.alloc_fragment([block_h], accum_dtype) + row_sum = T.alloc_fragment([block_h], accum_dtype) + tile_sum = T.alloc_fragment([block_h], accum_dtype) + alpha = T.alloc_fragment([block_h], accum_dtype) + + for item in T.Parallel(topk): + Target[row, item] = 0 + + for head_tile in T.serial(num_head_tiles): + for head, d in T.Parallel(block_h, dim): + head_index = head_tile * block_h + head + query_shared[head, d] = T.if_then_else( + head_index < heads, Query[row, head_index, d], 0 + ) + T.fill(row_max, -(2**30)) + T.fill(row_sum, 0) + + for tile in T.Pipelined(num_tiles, num_stages=num_stages): + for item in T.Parallel(block_i): + index = Indices[row, tile * block_i + item] + valid[item] = index >= 0 and index < seq_len_kv + for item, d in T.Parallel(block_i, dim): + index = Indices[row, tile * block_i + item] + safe_index = T.max(T.min(index, seq_len_kv - 1), 0) + key_shared[item, d] = T.if_then_else(valid[item], Key[safe_index, d], 0) + T.gemm( + query_shared, + key_shared, + scores, + transpose_B=True, + clear_accum=True, + policy=T.GemmWarpPolicy.FullRow, + ) + for head, item in T.Parallel(block_h, block_i): + scores[head, item] = T.if_then_else( + valid[item] and head_tile * block_h + head < heads, + scores[head, item], + -T.infinity(accum_dtype), + ) + T.copy(row_max, previous_max) + T.reduce_max(scores, tile_max, dim=1, clear=True) + for head in T.Parallel(block_h): + row_max[head] = T.max(previous_max[head], tile_max[head]) + alpha[head] = T.exp2((previous_max[head] - row_max[head]) * scale_log2) + for head, item in T.Parallel(block_h, block_i): + probabilities[head, item] = T.if_then_else( + valid[item] and head_tile * block_h + head < heads, + T.exp2((scores[head, item] - row_max[head]) * scale_log2), + 0, + ) + T.reduce_sum(probabilities, tile_sum, dim=1, clear=True) + for head in T.Parallel(block_h): + row_sum[head] = row_sum[head] * alpha[head] + tile_sum[head] + + for tile in T.Pipelined(num_tiles, num_stages=num_stages): + for item in T.Parallel(block_i): + index = Indices[row, tile * block_i + item] + valid[item] = index >= 0 and index < seq_len_kv + for item, d in T.Parallel(block_i, dim): + index = Indices[row, tile * block_i + item] + safe_index = T.max(T.min(index, seq_len_kv - 1), 0) + key_shared[item, d] = T.if_then_else(valid[item], Key[safe_index, d], 0) + T.gemm( + query_shared, + key_shared, + scores, + transpose_B=True, + clear_accum=True, + policy=T.GemmWarpPolicy.FullRow, + ) + for head, item in T.Parallel(block_h, block_i): + scores[head, item] = T.if_then_else( + valid[item] and head_tile * block_h + head < heads, + scores[head, item], + -T.infinity(accum_dtype), + ) + probabilities[head, item] = T.if_then_else( + valid[item] + and head_tile * block_h + head < heads + and row_sum[head] > 0, + T.exp2((scores[head, item] - row_max[head]) * scale_log2) + / row_sum[head], + 0, + ) + T.reduce_sum(probabilities, target_tile, dim=0, clear=True) + for item in T.Parallel(block_i): + Target[row, tile * block_i + item] += target_tile[item] + + return main + + +@tilelang_jit(out_idx=[-2, -1], pass_configs=_PASS_CONFIGS) +def sparse_indexer_kl(topk: int, block_i: int = 256, threads: int = 256): # pragma: no cover + """Build a kernel that computes sparse KL row sums and gradients for indexer logits.""" + require_tilelang() + assert topk % block_i == 0 + + seq_len = T.dynamic("seq_len") + accum_dtype = T.float32 + num_tiles = tilelang.cdiv(topk, block_i) + log2_e = 1.4426950408889634 + ln_2 = 0.6931471805599453 + eps = 1.0e-10 + + @T.prim_func + def main( + Target: T.Tensor([seq_len, topk], accum_dtype), # type: ignore + IndexLogits: T.Tensor([seq_len, topk], accum_dtype), # type: ignore + ValidMask: T.Tensor([seq_len, topk], "bool"), # type: ignore + GradLogits: T.Tensor([seq_len, topk], accum_dtype), # type: ignore + KLRows: T.Tensor([seq_len], accum_dtype), # type: ignore + ): + with T.Kernel(seq_len, threads=threads) as row: + logits = T.alloc_fragment([1, block_i], accum_dtype) + target = T.alloc_fragment([1, block_i], accum_dtype) + probabilities = T.alloc_fragment([1, block_i], accum_dtype) + kl_terms = T.alloc_fragment([1, block_i], accum_dtype) + valid = T.alloc_fragment([block_i], "bool") + row_max = T.alloc_fragment([1], accum_dtype) + previous_max = T.alloc_fragment([1], accum_dtype) + tile_max = T.alloc_fragment([1], accum_dtype) + row_sum = T.alloc_fragment([1], accum_dtype) + tile_sum = T.alloc_fragment([1], accum_dtype) + target_sum = T.alloc_fragment([1], accum_dtype) + target_tile_sum = T.alloc_fragment([1], accum_dtype) + kl_sum = T.alloc_fragment([1], accum_dtype) + kl_tile_sum = T.alloc_fragment([1], accum_dtype) + + T.fill(row_max, -(2**30)) + T.fill(row_sum, 0) + T.fill(target_sum, 0) + T.fill(kl_sum, 0) + + for tile in T.serial(num_tiles): + for item in T.Parallel(block_i): + valid[item] = ValidMask[row, tile * block_i + item] + logits[0, item] = T.if_then_else( + valid[item], + IndexLogits[row, tile * block_i + item], + -T.infinity(accum_dtype), + ) + target[0, item] = T.if_then_else( + valid[item], Target[row, tile * block_i + item], 0 + ) + T.copy(row_max, previous_max) + T.reduce_max(logits, tile_max, dim=1, clear=True) + row_max[0] = T.max(previous_max[0], tile_max[0]) + for item in T.Parallel(block_i): + probabilities[0, item] = T.if_then_else( + valid[item], T.exp2((logits[0, item] - row_max[0]) * log2_e), 0 + ) + T.reduce_sum(probabilities, tile_sum, dim=1, clear=True) + row_sum[0] = ( + row_sum[0] * T.exp2((previous_max[0] - row_max[0]) * log2_e) + tile_sum[0] + ) + T.reduce_sum(target, target_tile_sum, dim=1, clear=True) + target_sum[0] += target_tile_sum[0] + + for tile in T.serial(num_tiles): + for item in T.Parallel(block_i): + valid[item] = ValidMask[row, tile * block_i + item] + logits[0, item] = T.if_then_else( + valid[item], + IndexLogits[row, tile * block_i + item], + -T.infinity(accum_dtype), + ) + target[0, item] = T.if_then_else( + valid[item] and target_sum[0] > 0, + Target[row, tile * block_i + item] / target_sum[0], + 0, + ) + probabilities[0, item] = T.if_then_else( + valid[item] and row_sum[0] > 0, + T.exp2((logits[0, item] - row_max[0]) * log2_e) / row_sum[0], + 0, + ) + GradLogits[row, tile * block_i + item] = T.if_then_else( + valid[item], probabilities[0, item] - target[0, item], 0 + ) + kl_terms[0, item] = T.if_then_else( + valid[item] and target[0, item] > 0, + target[0, item] + * ( + T.log2(T.max(target[0, item], eps)) * ln_2 + - (logits[0, item] - row_max[0]) + + T.log2(T.max(row_sum[0], eps)) * ln_2 + ), + 0, + ) + T.reduce_sum(kl_terms, kl_tile_sum, dim=1, clear=True) + kl_sum[0] += kl_tile_sum[0] + + KLRows[row] = kl_sum[0] + + return main + + +def sparse_indexer_target_interface( + query: torch.Tensor, key: torch.Tensor, topk_indices: torch.Tensor, softmax_scale: float +) -> torch.Tensor: + """Compute the local-head sparse attention target on selected top-k keys.""" + require_tilelang() + seq_len, heads, dim = query.shape + topk = topk_indices.size(1) + kernel = _get_target_kernel(heads, dim, topk, softmax_scale) + return kernel(query, key, topk_indices) + + +def sparse_indexer_kl_interface( + target: torch.Tensor, index_logits: torch.Tensor, valid_mask: torch.Tensor +) -> tuple[torch.Tensor, torch.Tensor]: + """Compute unscaled indexer KL sum and its exact gradient with respect to logits.""" + require_tilelang() + kernel = _get_kl_kernel(valid_mask.size(1)) + grad_logits, kl_rows = kernel(target, index_logits, valid_mask) + return kl_rows.sum(), grad_logits + + +class SparseIndexerKLLoss(torch.autograd.Function): # pragma: no cover + """Autograd bridge from fused sparse KL score gradients to the TileLang indexer.""" + + @staticmethod + def forward(ctx, target, index_logits, valid_mask): + """Compute the sparse indexer KL loss and save its logits gradient.""" + kl_sum, grad_logits = sparse_indexer_kl_interface(target, index_logits, valid_mask) + ctx.save_for_backward(grad_logits) + return kl_sum + + @staticmethod + def backward(ctx, grad_output): + """Scale the saved index-logits gradient for the backward pass.""" + (grad_logits,) = ctx.saved_tensors + return None, grad_logits * grad_output, None + + +if not HAVE_TILELANG: + SparseIndexerKLLoss = None diff --git a/megatron/core/transformer/experimental_attention_variant/ops/tilelang_sparse_mla_bwd.py b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_sparse_mla_bwd.py new file mode 100644 index 00000000000..1cccec8339b --- /dev/null +++ b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_sparse_mla_bwd.py @@ -0,0 +1,529 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# ruff: noqa +# Adapted from: +# https://github.com/tile-ai/tilelang/blob/4ff81c7d40803d269569e157e847623e84553f78/ +# examples/deepseek_v32/sparse_mla_bwd.py +import threading +from collections import OrderedDict + +import torch + +from .tilelang_utils import ( + HAVE_TILELANG, + T, + _env_int, + _get_cached_kernel, + _normalize_sm_scale, + _round_up, + require_tilelang, + tilelang, + tilelang_jit, +) + +_SPARSE_MLA_BWD_BLOCK_SIZE = 32 +_tilelang_sparse_mla_preprocess_kernel_cache = OrderedDict() +_tilelang_sparse_mla_bwd_kernel_cache = OrderedDict() +_tilelang_sparse_mla_postprocess_kernel_cache = OrderedDict() +_tilelang_sparse_mla_bwd_cache_lock = threading.Lock() + + +def _get_preprocess_kernel(H: int, D: int): + return _get_cached_kernel( + _tilelang_sparse_mla_preprocess_kernel_cache, + _tilelang_sparse_mla_bwd_cache_lock, + (H, D), + lambda: preprocess(H, D), + ) + + +def _normalize_block_h(block_h: int) -> int: + if block_h >= 64: + return 64 + if block_h >= 32: + return 32 + return 16 + + +def _get_bwd_kernel( + H: int, D: int, D_tail: int, topk: int, kv_group: int, sm_scale, max_block_h: int +): + max_block_h = _normalize_block_h(max_block_h) + key = (H, D, D_tail, topk, kv_group, _normalize_sm_scale(sm_scale), max_block_h) + return _get_cached_kernel( + _tilelang_sparse_mla_bwd_kernel_cache, + _tilelang_sparse_mla_bwd_cache_lock, + key, + lambda: bwd(H, D, D_tail, topk, kv_group, sm_scale, max_block_h=max_block_h), + ) + + +def _get_postprocess_kernel(D: int, D_tail: int, kv_group: int): + return _get_cached_kernel( + _tilelang_sparse_mla_postprocess_kernel_cache, + _tilelang_sparse_mla_bwd_cache_lock, + (D, D_tail, kv_group), + lambda: postprocess(D, D_tail, kv_group), + ) + + +@tilelang_jit(out_idx=[-1]) +def preprocess( # pragma: no cover + H, + D, + block_ND=32, + num_stages=5, + dtype=T.bfloat16 if HAVE_TILELANG else None, + accum_dtype=T.float32 if HAVE_TILELANG else None, +): + """Build preprocessing kernel that computes Delta = sum(O * dO) per row/head.""" + require_tilelang() + assert dtype == T.bfloat16 + assert accum_dtype == T.float32 + batch = T.dynamic("batch") + seq_len = T.dynamic("seq_len") + shape = [batch, seq_len, H, D] + + @T.prim_func + def preprocess_kernel( + O: T.Tensor(shape, dtype), + dO: T.Tensor(shape, dtype), + Delta: T.Tensor([batch, seq_len, H], accum_dtype), + ): + with T.Kernel(H, T.ceildiv(seq_len, block_ND), batch) as (bx, by, bz): + o = T.alloc_fragment([block_ND, block_ND], accum_dtype) + do = T.alloc_fragment([block_ND, block_ND], accum_dtype) + delta = T.alloc_fragment([block_ND], accum_dtype) + acc = T.alloc_fragment([block_ND, block_ND], accum_dtype) + T.clear(acc) + for k in T.Pipelined(T.ceildiv(D, block_ND), num_stages=num_stages): + T.copy( + O[ + bz, + by * block_ND : (by + 1) * block_ND, + bx, + k * block_ND : (k + 1) * block_ND, + ], + o, + ) + T.copy( + dO[ + bz, + by * block_ND : (by + 1) * block_ND, + bx, + k * block_ND : (k + 1) * block_ND, + ], + do, + ) + for i, j in T.Parallel(block_ND, block_ND): + acc[i, j] += o[i, j] * do[i, j] + T.reduce_sum(acc, delta, 1) + T.copy(delta, Delta[bz, by * block_ND : (by + 1) * block_ND, bx]) + + return preprocess_kernel + + +@tilelang_jit(out_idx=[-1]) +def postprocess( # pragma: no cover + D, + D_tail, + kv_group=1, + block_N=64, + threads=128, + dtype=T.bfloat16 if HAVE_TILELANG else None, + accum_dtype=T.float32 if HAVE_TILELANG else None, +): + """Build postprocess kernel that casts/exports accumulated dKV.""" + require_tilelang() + assert dtype == T.bfloat16 + assert accum_dtype == T.float32 + batch = T.dynamic("batch") + seq_len_kv = T.dynamic("seq_len_kv") + dkv_shape = [batch, seq_len_kv, kv_group, D + D_tail] + + @T.prim_func + def postprocess_kernel( + dKV: T.Tensor(dkv_shape, accum_dtype), dKV_out: T.Tensor(dkv_shape, dtype) + ): + with T.Kernel(T.ceildiv(seq_len_kv, block_N), kv_group, batch, threads=threads) as ( + bx, + by, + bz, + ): + T.copy( + dKV[bz, bx * block_N : (bx + 1) * block_N, by, :], + dKV_out[bz, bx * block_N : (bx + 1) * block_N, by, :], + ) + + return postprocess_kernel + + +_SPARSE_MLA_BWD_PASS_CONFIGS = ( + { + tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, + tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, + tilelang.PassConfigKey.TL_ENABLE_AGGRESSIVE_SHARED_MEMORY_MERGE: True, + } + if HAVE_TILELANG + else {} +) + + +@tilelang_jit(out_idx=[-2], pass_configs=_SPARSE_MLA_BWD_PASS_CONFIGS) +def bwd( # pragma: no cover + H, + D, + D_tail, + topk, + kv_group=1, + sm_scale=None, + block_size=32, + max_block_h=32, + num_stages=2, + threads=128, + indices_dtype=T.int32 if HAVE_TILELANG else None, + dtype=T.bfloat16 if HAVE_TILELANG else None, + accum_dtype=T.float32 if HAVE_TILELANG else None, +): + """Build sparse-MLA backward kernel.""" + require_tilelang() + assert ( + topk % block_size == 0 + ), "otherwise will load some index=0 thus causing wrong kv to be loaded" + assert dtype == T.bfloat16 + assert accum_dtype == T.float32 + assert indices_dtype == T.int32 + + if sm_scale is None: + sm_scale = (D + D_tail) ** (-0.5) + sm_scale_mul_reciprocal_log2 = sm_scale * 1.44269504 # log2(e) + + batch = T.dynamic("batch") + seq_len = T.dynamic("seq_len") + seq_len_kv = T.dynamic("seq_len_kv") + + H_kv = H // kv_group + q_shape = [batch, seq_len, H, D + D_tail] + k_shape = [batch, seq_len_kv, kv_group, D + D_tail] + o_shape = [batch, seq_len, H, D] + indices_shape = [batch, seq_len, kv_group, topk] + delta_shape = [batch, seq_len, H] + lse_shape = [batch, seq_len, H] + assert indices_dtype == T.int32 + assert dtype == T.bfloat16 + assert accum_dtype == T.float32 + + H = H_kv + padded_H = max(tilelang.math.next_power_of_2(H_kv), 16) + block_H = min(_normalize_block_h(max_block_h), padded_H) + assert padded_H % block_H == 0 + NH = padded_H // block_H + BS = block_size + NS = tilelang.cdiv(topk, block_size) + + split_store = 2 + + @T.prim_func + def sparse_mla_bwd_kernel( + Q: T.Tensor(q_shape, dtype), + KV: T.Tensor(k_shape, dtype), + dO: T.Tensor(o_shape, dtype), + Indices: T.Tensor(indices_shape, indices_dtype), + Lse: T.Tensor(lse_shape, accum_dtype), + Delta: T.Tensor(delta_shape, accum_dtype), + dQ: T.Tensor(q_shape, dtype), + dKV: T.Tensor(k_shape, accum_dtype), + ): + with T.Kernel(seq_len, batch, kv_group * NH, threads=threads) as (s_i, by, bz): + Q_shared = T.alloc_shared([block_H, D], dtype) + Q_tail_shared = T.alloc_shared([block_H, D_tail], dtype) + KV_shared = T.alloc_shared([BS, D], dtype) + KV_tail_shared = T.alloc_shared([BS, D_tail], dtype) + dO_shared = T.alloc_shared([block_H, D], dtype) + mask = T.alloc_fragment([BS], "bool") + + P_shared_cast = T.alloc_shared([block_H, BS], dtype) + dP_shared_cast = T.alloc_shared([block_H, BS], dtype) + dQ_shared = T.alloc_shared([block_H, D], dtype) + dQ_tail_shared = T.alloc_shared([block_H, D_tail], dtype) + + acc_p = T.alloc_fragment([block_H, BS], accum_dtype) + acc_dp = T.alloc_fragment([block_H, BS], accum_dtype) + acc_dq = T.alloc_fragment([block_H, D], accum_dtype) + acc_dq_tail = T.alloc_fragment([block_H, D_tail], accum_dtype) + acc_dkv = T.alloc_fragment([BS, D], accum_dtype) + acc_dkv_tail = T.alloc_fragment([BS, D_tail], accum_dtype) + acc_dkv_shared = T.alloc_shared([BS // split_store, D], accum_dtype) + acc_dkv_tail_shared = T.alloc_shared([BS // split_store, D_tail], accum_dtype) + + T.copy(Q[by, s_i, bz * block_H : (bz + 1) * block_H, :D], Q_shared) + T.copy(Q[by, s_i, bz * block_H : (bz + 1) * block_H, D:], Q_tail_shared) + T.copy(dO[by, s_i, bz * block_H : (bz + 1) * block_H, :D], dO_shared) + + T.clear(acc_dq) + T.clear(acc_dq_tail) + + # Process each block of indices + for i_i in T.Pipelined(NS, num_stages=num_stages): + # Check which indices are valid + for bi_i in T.Parallel(BS): + # Changed here for thd + mask[bi_i] = Indices[by, s_i, bz // NH, i_i * BS + bi_i] != -1 + + # Compute attention scores + for h_i, bi_i in T.Parallel(block_H, BS): + acc_p[h_i, bi_i] = T.if_then_else(mask[bi_i], 0, -T.infinity(acc_p.dtype)) + + # Load KV, V for this block of indices + for bi_i, d_i in T.Parallel(BS, D): + idx = Indices[by, s_i, bz // NH, i_i * BS + bi_i] + safe_idx = T.max(idx, 0) + KV_shared[bi_i, d_i] = KV[by, safe_idx, bz // NH, d_i] + + T.gemm( + Q_shared, KV_shared, acc_p, transpose_B=True, policy=T.GemmWarpPolicy.FullCol + ) + + for bi_i, d_i in T.Parallel(BS, D_tail): + idx = Indices[by, s_i, bz // NH, i_i * BS + bi_i] + safe_idx = T.max(idx, 0) + KV_tail_shared[bi_i, d_i] = KV[by, safe_idx, bz // NH, D + d_i] + T.gemm( + Q_tail_shared, + KV_tail_shared, + acc_p, + transpose_B=True, + policy=T.GemmWarpPolicy.FullCol, + ) + + for h_i, bi_i in T.Parallel(block_H, BS): + acc_p[h_i, bi_i] = T.exp2( + acc_p[h_i, bi_i] * sm_scale_mul_reciprocal_log2 + - Lse[by, s_i, bz * block_H + h_i] + ) + + T.copy(acc_p, P_shared_cast) + + T.gemm( + dO_shared, + KV_shared, + acc_dp, + transpose_B=True, + policy=T.GemmWarpPolicy.FullCol, + clear_accum=True, + ) + + for h_i, bi_i in T.Parallel(block_H, BS): + acc_dp[h_i, bi_i] = ( + acc_p[h_i, bi_i] + * (acc_dp[h_i, bi_i] - Delta[by, s_i, bz * block_H + h_i]) + * sm_scale + ) + + T.copy(acc_dp, dP_shared_cast) + T.gemm(dP_shared_cast, KV_shared, acc_dq, policy=T.GemmWarpPolicy.FullCol) + T.gemm(dP_shared_cast, KV_tail_shared, acc_dq_tail, policy=T.GemmWarpPolicy.FullCol) + + T.gemm( + dP_shared_cast, + Q_shared, + acc_dkv, + transpose_A=True, + policy=T.GemmWarpPolicy.FullCol, + clear_accum=True, + ) + T.gemm( + P_shared_cast, + dO_shared, + acc_dkv, + transpose_A=True, + policy=T.GemmWarpPolicy.FullCol, + ) + + T.clear(acc_dkv_tail) + T.gemm( + dP_shared_cast, + Q_tail_shared, + acc_dkv_tail, + transpose_A=True, + policy=T.GemmWarpPolicy.FullCol, + ) + + for s in range(split_store): + for bi_i, d_i in T.Parallel(BS, D): + if bi_i < BS // split_store: + acc_dkv_shared[bi_i, d_i] = acc_dkv[bi_i + s * (BS // split_store), d_i] + + for bi_i, d_i in T.Parallel(BS, D_tail): + if bi_i < BS // split_store: + acc_dkv_tail_shared[bi_i, d_i] = acc_dkv_tail[ + bi_i + s * (BS // split_store), d_i + ] + + for bi_i, d_i in T.Parallel(BS // split_store, D // 4): + idx = Indices[by, s_i, bz // NH, i_i * BS + bi_i + s * (BS // split_store)] + if idx >= 0: + T.atomic_addx4( + dKV[by, idx, bz // NH, d_i * 4], acc_dkv_shared[bi_i, d_i * 4] + ) + + # Atomically update dKV, dKV_tail tensors + for bi_i, d_i in T.Parallel(BS // split_store, D_tail // 4): + idx = Indices[by, s_i, bz // NH, i_i * BS + bi_i + s * (BS // split_store)] + if idx >= 0: + T.atomic_addx4( + dKV[by, idx, bz // NH, D + d_i * 4], + acc_dkv_tail_shared[bi_i, d_i * 4], + ) + + # Store the accumulated dQ + T.copy(acc_dq, dQ_shared) + T.copy(acc_dq_tail, dQ_tail_shared) + + T.copy(dQ_shared, dQ[by, s_i, bz * block_H : (bz + 1) * block_H, :D]) + T.copy(dQ_tail_shared, dQ[by, s_i, bz * block_H : (bz + 1) * block_H, D:]) + + return sparse_mla_bwd_kernel + + +def _sparse_mla_delta_batched(o, do): # pragma: no cover + """Compute Delta = sum(O * dO) with safe sequence padding for TileLang tiles.""" + require_tilelang() + assert o.is_contiguous() + assert do.is_contiguous() + assert o.shape == do.shape + B, S, H, D = o.shape + + seq_len_padded = _round_up(S, _SPARSE_MLA_BWD_BLOCK_SIZE) + if seq_len_padded != S: + o_padded = torch.zeros((B, seq_len_padded, H, D), dtype=o.dtype, device=o.device) + o_padded[:, :S].copy_(o) + o = o_padded + + do_padded = torch.zeros((B, seq_len_padded, H, D), dtype=do.dtype, device=do.device) + do_padded[:, :S].copy_(do) + do = do_padded + + preprocess_kernel = _get_preprocess_kernel(H, D) + return preprocess_kernel(o, do)[:, :S].contiguous() + + +def sparse_mla_delta(o, do): # pragma: no cover + """Compute Delta = sum(O * dO) per sequence row and head.""" + squeeze_batch = o.ndim == 3 + if squeeze_batch: + o = o.unsqueeze(0) + do = do.unsqueeze(0) + delta = _sparse_mla_delta_batched(o, do) + if squeeze_batch: + delta = delta.squeeze(0) + return delta + + +def sparse_mla_bwd(q, kv, o, do, indices, lse, sm_scale=None, delta=None): # pragma: no cover + """Run sparse-MLA backward kernels and return (dq, dkv).""" + require_tilelang() + + seq_bucket = _env_int("MCORE_DSA_TILELANG_SEQ_BUCKET", 256) + topk_bucket = _env_int("MCORE_DSA_TILELANG_TOPK_BUCKET", _SPARSE_MLA_BWD_BLOCK_SIZE) + max_block_h = _env_int("MCORE_DSA_TILELANG_BWD_MAX_BLOCK_H", 32) + + squeeze_batch = q.ndim == 3 + if squeeze_batch: + q = q.unsqueeze(0) + kv = kv.unsqueeze(0) + do = do.unsqueeze(0) + indices = indices.unsqueeze(0) + lse = lse.unsqueeze(0) + if o is not None: + if squeeze_batch: + o = o.unsqueeze(0) + + assert q.is_contiguous() + assert kv.is_contiguous() + assert indices.is_contiguous() + assert lse.is_contiguous() + assert q.ndim == 4 and kv.ndim == 4 and do.ndim == 4 and indices.ndim == 4 and lse.ndim == 3 + B, S, H, dim_plus_tail_dim = q.shape + _, S_kv, kv_group, _ = kv.shape + assert kv.shape[-1] == dim_plus_tail_dim + assert kv.shape[0] == B + # This copied kernel currently assumes a fixed base value-channel dimension. + D = 512 + assert ( + dim_plus_tail_dim >= D + ), f"Invalid dimensions: dim_plus_tail_dim={dim_plus_tail_dim} is smaller than base D={D}" + + D_tail = dim_plus_tail_dim - D + topk = indices.shape[-1] + assert indices.shape == (B, S, kv_group, topk) + assert lse.shape == (B, S, H) + + seq_bucketed = _round_up(S, seq_bucket) + seq_kv_bucketed = _round_up(S_kv, seq_bucket) + topk_bucketed = _round_up(_round_up(topk, topk_bucket), _SPARSE_MLA_BWD_BLOCK_SIZE) + + if seq_bucketed != S: + q_padded = torch.zeros( + (B, seq_bucketed, H, dim_plus_tail_dim), dtype=q.dtype, device=q.device + ) + q_padded[:, :S].copy_(q) + q = q_padded + + if o is not None: + o_padded = torch.zeros((B, seq_bucketed, H, D), dtype=o.dtype, device=o.device) + o_padded[:, :S].copy_(o) + o = o_padded + + do_padded = torch.zeros((B, seq_bucketed, H, D), dtype=do.dtype, device=do.device) + do_padded[:, :S].copy_(do) + do = do_padded + + lse_padded = torch.zeros((B, seq_bucketed, H), dtype=lse.dtype, device=lse.device) + lse_padded[:, :S].copy_(lse) + lse = lse_padded + + if seq_kv_bucketed != S_kv: + kv_padded = torch.zeros( + (B, seq_kv_bucketed, kv_group, dim_plus_tail_dim), dtype=kv.dtype, device=kv.device + ) + kv_padded[:, :S_kv].copy_(kv) + kv = kv_padded + + if seq_bucketed != S or topk_bucketed != topk: + indices_padded = torch.full( + (B, seq_bucketed, kv_group, topk_bucketed), + -1, + dtype=indices.dtype, + device=indices.device, + ) + indices_padded[:, :S, :, :topk].copy_(indices) + indices = indices_padded + + if delta is not None: + if delta.ndim == 2: + delta = delta.unsqueeze(0) + if seq_bucketed != S: + delta_padded = torch.zeros((B, seq_bucketed, H), dtype=delta.dtype, device=delta.device) + delta_padded[:, :S].copy_(delta) + delta = delta_padded + + # Get kernels + bwd_kernel = _get_bwd_kernel(H, D, D_tail, topk_bucketed, kv_group, sm_scale, max_block_h) + postprocess_kernel = _get_postprocess_kernel(D, D_tail, kv_group) + + if delta is None: + if o is None: + raise ValueError("sparse_mla_bwd requires either output tensor o or precomputed delta") + delta = _sparse_mla_delta_batched(o, do) + dkv = torch.zeros_like(kv, dtype=torch.float32) + dq = bwd_kernel(q, kv, do, indices, lse, delta, dkv) + dkv = postprocess_kernel(dkv) + + dq = dq[:, :S].contiguous() + dkv = dkv[:, :S_kv].contiguous() + + if squeeze_batch: + dq = dq.squeeze(0) + dkv = dkv.squeeze(0) + + return dq, dkv diff --git a/megatron/core/transformer/experimental_attention_variant/ops/tilelang_sparse_mla_fwd.py b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_sparse_mla_fwd.py new file mode 100644 index 00000000000..707c453d00e --- /dev/null +++ b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_sparse_mla_fwd.py @@ -0,0 +1,310 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# ruff: noqa +# Adapted from: +# https://github.com/tile-ai/tilelang/blob/e666d2d3cc483829c57618c9ebf2e4f4ada0819d/ +# examples/deepseek_v32/sparse_mla_fwd.py +import threading +from collections import OrderedDict + +import torch + +from .tilelang_utils import ( + HAVE_TILELANG, + T, + _env_int, + _get_cached_kernel, + _normalize_sm_scale, + _round_up, + require_tilelang, + tilelang, + tilelang_jit, +) + +_tilelang_sparse_mla_fwd_kernel_cache = OrderedDict() +_tilelang_sparse_mla_fwd_cache_lock = threading.Lock() + + +def _get_sparse_mla_fwd_kernel( + heads: int, + dim: int, + tail_dim: int, + topk: int, + kv_group: int, + sm_scale, + block_I: int, + num_stages: int, + threads: int, +): + key = ( + heads, + dim, + tail_dim, + topk, + kv_group, + _normalize_sm_scale(sm_scale), + block_I, + num_stages, + threads, + ) + return _get_cached_kernel( + _tilelang_sparse_mla_fwd_kernel_cache, + _tilelang_sparse_mla_fwd_cache_lock, + key, + lambda: sparse_mla_fwd( + heads, + dim, + tail_dim, + topk, + kv_group, + sm_scale, + block_I=block_I, + num_stages=num_stages, + threads=threads, + ), + ) + + +_SPARSE_MLA_FWD_PASS_CONFIGS = ( + { + tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, + tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, + } + if HAVE_TILELANG + else {} +) + + +@tilelang_jit(out_idx=[-2, -1], pass_configs=_SPARSE_MLA_FWD_PASS_CONFIGS) +def sparse_mla_fwd( # pragma: no cover + heads, dim, tail_dim, topk, kv_group=1, sm_scale=None, block_I=64, num_stages=2, threads=256 +): + """Build sparse-MLA forward kernel.""" + require_tilelang() + assert dim == tilelang.math.next_power_of_2(dim), f"dim must be a power of two, got dim={dim}" + assert tail_dim == tilelang.math.next_power_of_2( + tail_dim + ), f"tail_dim must be a power of two, got tail_dim={tail_dim}" + assert ( + topk % block_I == 0 + ), "otherwise will load some index=0 thus causing wrong kv to be loaded" + if sm_scale is None: + sm_scale = (1.0 / (dim + tail_dim)) ** 0.5 * 1.44269504 # log2(e) + else: + sm_scale = sm_scale * 1.44269504 # log2(e) + + batch = T.dynamic("batch") + seq_len = T.dynamic("seq_len") + seq_len_kv = T.dynamic("seq_len_kv") + + head_kv = heads // kv_group + q_shape = [batch, seq_len, heads, dim + tail_dim] + kv_shape = [batch, seq_len_kv, kv_group, dim + tail_dim] + o_shape = [batch, seq_len, heads, dim] + indices_shape = [batch, seq_len, kv_group, topk] + lse_shape = [batch, seq_len, heads] + indices_dtype = T.int32 + dtype = T.bfloat16 + accum_dtype = T.float32 + + G = kv_group + H = head_kv + padded_H = max(tilelang.math.next_power_of_2(head_kv), 16) + if padded_H != H: + assert kv_group == 1, ( + "here we solve the H padding automatically, otherwise handle Q/Output copy with " + "your own mask (for kv_group==1, g_i*padded_H:(g_i+1)*padded_H is handled)" + ) + BI = block_I + NI = tilelang.cdiv(topk, block_I) + D = dim + D_tail = tail_dim + + if head_kv > 64: + assert head_kv % 64 == 0, "head_kv should be a multiple of 64" + REPLICATE_H = head_kv // 64 + else: + REPLICATE_H = 1 + + H_per_block = padded_H if REPLICATE_H == 1 else 64 + + @T.prim_func + def main( + Q: T.Tensor(q_shape, dtype), # type: ignore + KV: T.Tensor(kv_shape, dtype), # type: ignore + Indices: T.Tensor(indices_shape, indices_dtype), # type: ignore + Output: T.Tensor(o_shape, dtype), # type: ignore + Lse: T.Tensor(lse_shape, accum_dtype), # type: ignore + ): + with T.Kernel(seq_len * REPLICATE_H, batch, kv_group, threads=threads) as (bx, by, bz): + Q_shared = T.alloc_shared([H_per_block, D], dtype) + Q_tail_shared = T.alloc_shared([H_per_block, D_tail], dtype) + KV_shared = T.alloc_shared([BI, D], dtype) + K_tail_shared = T.alloc_shared([BI, D_tail], dtype) + O_shared = T.alloc_shared([H_per_block, D], dtype) + Lse_shared = T.alloc_shared([H_per_block], accum_dtype) + mask = T.alloc_fragment([BI], "bool") + + acc_o = T.alloc_fragment([H_per_block, D], accum_dtype) + acc_s = T.alloc_fragment([H_per_block, BI], accum_dtype) + S_shared = T.alloc_shared([H_per_block, BI], dtype) + sumexp = T.alloc_fragment([H_per_block], accum_dtype) + sumexp_i = T.alloc_fragment([H_per_block], accum_dtype) + alpha = T.alloc_fragment([H_per_block], accum_dtype) + m_i = T.alloc_fragment([H_per_block], accum_dtype) + m_i_prev = T.alloc_fragment([H_per_block], accum_dtype) + + T.fill(acc_o, 0) + T.fill(sumexp, 0) + T.fill(m_i, -(2**30)) # avoid -inf - inf to cause nan + + b_i, g_i = by, bz + s_i = bx if REPLICATE_H == 1 else (bx // REPLICATE_H) + q_i = s_i + max_kv_i = q_i + + H0 = g_i * padded_H + (0 if REPLICATE_H == 1 else (bx % REPLICATE_H) * 64) + H1 = H0 + H_per_block + + T.copy(Q[b_i, s_i, H0:H1, :D], Q_shared) + T.copy(Q[b_i, s_i, H0:H1, D:], Q_tail_shared) + + for i_i in T.Pipelined(NI, num_stages=num_stages): + for bi_i in T.Parallel(BI): + # Changed here for thd + mask[bi_i] = Indices[b_i, s_i, g_i, i_i * BI + bi_i] != -1 + + for bi_i, d_i in T.Parallel(BI, D): + idx = Indices[b_i, s_i, g_i, i_i * BI + bi_i] + safe_idx = T.max(idx, 0) + KV_shared[bi_i, d_i] = KV[b_i, safe_idx, g_i, d_i] + for bi_i, d_i in T.Parallel(BI, D_tail): + idx = Indices[b_i, s_i, g_i, i_i * BI + bi_i] + safe_idx = T.max(idx, 0) + K_tail_shared[bi_i, d_i] = KV[b_i, safe_idx, g_i, D + d_i] + + for h_i, bi_i in T.Parallel(H_per_block, BI): + acc_s[h_i, bi_i] = T.if_then_else(mask[bi_i], 0, -T.infinity(acc_s.dtype)) + T.gemm( + Q_shared, KV_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullRow + ) + T.gemm( + Q_tail_shared, + K_tail_shared, + acc_s, + transpose_B=True, + policy=T.GemmWarpPolicy.FullRow, + ) + T.copy(m_i, m_i_prev) + T.reduce_max(acc_s, m_i, dim=1, clear=False) + for h_i in T.Parallel(H_per_block): + m_i[h_i] = T.max(m_i[h_i], m_i_prev[h_i]) + for h_i in T.Parallel(H_per_block): + alpha[h_i] = T.exp2((m_i_prev[h_i] - m_i[h_i]) * sm_scale) + for h_i, bi_i in T.Parallel(H_per_block, BI): + acc_s[h_i, bi_i] = T.exp2(acc_s[h_i, bi_i] * sm_scale - m_i[h_i] * sm_scale) + # Reduce the current tile; the online softmax accumulation happens below. + T.reduce_sum(acc_s, sumexp_i, dim=1) + for h_i in T.Parallel(H_per_block): + sumexp[h_i] = sumexp[h_i] * alpha[h_i] + sumexp_i[h_i] + for h_i, d_i in T.Parallel(H_per_block, D): + acc_o[h_i, d_i] = acc_o[h_i, d_i] * alpha[h_i] + + T.copy(acc_s, S_shared) + T.gemm(S_shared, KV_shared, acc_o, policy=T.GemmWarpPolicy.FullRow) + + # Rescale. Packed THD can produce sentinel-only rows; define those rows as zero + # output/LSE instead of dividing by a zero softmax denominator. + for h_i, d_i in T.Parallel(H_per_block, D): + acc_o[h_i, d_i] = T.if_then_else(sumexp[h_i] > 0, acc_o[h_i, d_i] / sumexp[h_i], 0) + for h_i in T.Parallel(H_per_block): + sumexp[h_i] = T.if_then_else( + sumexp[h_i] > 0, T.log2(sumexp[h_i]) + m_i[h_i] * sm_scale, 0 + ) + + T.copy(acc_o, Output[b_i, s_i, H0:H1, :]) + T.copy(sumexp, Lse[b_i, s_i, H0:H1]) + + return main + + +def sparse_mla_fwd_interface( + q, kv, indices, sm_scale=None, d_v=512, block_I=64, num_stages=2, threads=256 +): + """Run sparse-MLA forward kernel and return (out, lse).""" + require_tilelang() + seq_bucket = _env_int("MCORE_DSA_TILELANG_SEQ_BUCKET", 256) + topk_bucket = _env_int("MCORE_DSA_TILELANG_TOPK_BUCKET", block_I) + + squeeze_batch = q.ndim == 3 + if squeeze_batch: + q = q.unsqueeze(0) + kv = kv.unsqueeze(0) + indices = indices.unsqueeze(0) + + assert q.is_contiguous() and kv.is_contiguous() and indices.is_contiguous() + assert q.ndim == 4 and kv.ndim == 4 and indices.ndim == 4 + batch, seq_len, heads, dim_plus_tail_dim = q.shape + _, seq_len_kv, kv_group, kv_dim = kv.shape + assert ( + kv_dim == dim_plus_tail_dim + ), "q and kv must have the same embedding dimension on the last axis" + assert ( + dim_plus_tail_dim == 576 + ), "TileLang sparse MLA fwd is currently specialized for dim_plus_tail_dim=576" + dim = d_v + assert 0 < dim <= dim_plus_tail_dim, f"d_v must be in (0, {dim_plus_tail_dim}], but got {dim}" + + assert kv.shape[-1] == dim_plus_tail_dim + tail_dim = dim_plus_tail_dim - dim + assert kv.shape[0] == batch + _, _, _, topk = indices.shape + assert indices.shape == (batch, seq_len, kv_group, topk) + + seq_len_bucketed = _round_up(seq_len, seq_bucket) + seq_len_kv_bucketed = _round_up(seq_len_kv, seq_bucket) + topk_bucketed = _round_up(_round_up(topk, topk_bucket), block_I) + + if seq_len_bucketed != seq_len: + q_padded = torch.zeros( + (batch, seq_len_bucketed, heads, dim_plus_tail_dim), dtype=q.dtype, device=q.device + ) + q_padded[:, :seq_len].copy_(q) + q = q_padded + + if seq_len_kv_bucketed != seq_len_kv: + kv_padded = torch.zeros( + (batch, seq_len_kv_bucketed, kv_group, dim_plus_tail_dim), + dtype=kv.dtype, + device=kv.device, + ) + kv_padded[:, :seq_len_kv].copy_(kv) + kv = kv_padded + + if seq_len_bucketed != seq_len or topk_bucketed != topk: + indices_padded = torch.full( + (batch, seq_len_bucketed, kv_group, topk_bucketed), + -1, + dtype=indices.dtype, + device=indices.device, + ) + indices_padded[:, :seq_len, :, :topk].copy_(indices) + indices = indices_padded + + kernel = _get_sparse_mla_fwd_kernel( + heads=heads, + dim=dim, + tail_dim=tail_dim, + topk=topk_bucketed, + kv_group=kv_group, + sm_scale=sm_scale, + block_I=block_I, + num_stages=num_stages, + threads=threads, + ) + out, lse = kernel(q, kv, indices) + out = out[:, :seq_len].contiguous() + lse = lse[:, :seq_len].contiguous() + if squeeze_batch: + out = out.squeeze(0) + lse = lse.squeeze(0) + return out, lse diff --git a/megatron/core/transformer/experimental_attention_variant/ops/tilelang_utils.py b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_utils.py new file mode 100644 index 00000000000..689436bfb59 --- /dev/null +++ b/megatron/core/transformer/experimental_attention_variant/ops/tilelang_utils.py @@ -0,0 +1,101 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import os +from collections import OrderedDict + +import torch + +from megatron.core.utils import round_up_to_nearest_multiple + +try: + import tilelang + from tilelang import language as T # pylint: disable=unused-import + + HAVE_TILELANG = True +except (ImportError, OSError): + tilelang = None + T = None + HAVE_TILELANG = False + + +def _noop_jit(*args, **kwargs): + if len(args) == 1 and callable(args[0]) and not kwargs: + return args[0] + + def decorator(func): + return func + + return decorator + + +def tilelang_jit(*args, **kwargs): + """Return TileLang's jit decorator when available, otherwise a no-op decorator.""" + if HAVE_TILELANG: + return tilelang.jit(*args, **kwargs) + return _noop_jit(*args, **kwargs) + + +def require_tilelang(): + """Raise a clear error when a fused TileLang kernel is used without TileLang installed.""" + if not HAVE_TILELANG: + raise ImportError( + "TileLang is required to use fused DSA TileLang kernels. " + "Install tilelang or use the unfused fallback path." + ) + + +def _env_int(name: str, default: int) -> int: + """Parse a positive integer environment variable, falling back to ``default``.""" + value = os.getenv(name) + if value is None: + return default + try: + parsed = int(value) + except ValueError: + return default + return parsed if parsed > 0 else default + + +_TILELANG_KERNEL_CACHE_MAX = _env_int("MCORE_DSA_TILELANG_KERNEL_CACHE_MAX", 512) + + +def _cache_put_lru(cache: OrderedDict, key, value): + """Insert ``value`` as the most-recently-used entry, evicting oldest past the cap.""" + cache[key] = value + cache.move_to_end(key) + while len(cache) > _TILELANG_KERNEL_CACHE_MAX: + cache.popitem(last=False) + + +def _get_cached_kernel(cache: OrderedDict, lock, key, build_fn): + """Return a cached compiled kernel for ``key``, building it via ``build_fn`` on miss.""" + with lock: + kernel = cache.pop(key, None) + if kernel is None: + kernel = build_fn() + _cache_put_lru(cache, key, kernel) + return kernel + + +def _round_up(x: int, multiple: int) -> int: + if multiple <= 1: + return x + return round_up_to_nearest_multiple(x, multiple) + + +def _next_power_of_two(x: int) -> int: + if x <= 1: + return 1 + return 1 << (x - 1).bit_length() + + +def _normalize_sm_scale(sm_scale): + """Coerce a softmax scale to a stable float so it can key the kernel cache.""" + if sm_scale is None: + return None + if isinstance(sm_scale, torch.Tensor): + sm_scale = float(sm_scale.detach().item()) + else: + sm_scale = float(sm_scale) + # Avoid tiny floating-point jitter creating cache-key churn. + return round(sm_scale, 12) diff --git a/megatron/core/transformer/mla_qk_norm_config.py b/megatron/core/transformer/mla_qk_norm_config.py new file mode 100644 index 00000000000..e14066a105d --- /dev/null +++ b/megatron/core/transformer/mla_qk_norm_config.py @@ -0,0 +1,293 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +""" +Resolve MLA and DSA Q/KV norm configuration from a layer specification. +""" + +from typing import NoReturn + +from megatron.core.models.backends import get_backend +from megatron.core.transformer.identity_op import IdentityOp +from megatron.core.transformer.spec_utils import ModuleSpec +from megatron.core.transformer.torch_norm import LayerNormBuilder +from megatron.core.transformer.transformer_config import MLATransformerConfig + +__all__ = [] + +_QKNormResolvedConfig = dict[str, ModuleSpec | type | LayerNormBuilder] + + +class QKNormConfigResolver: + """Validate and resolve Q/KV norm placement for MLA and DSA. + + Q/KV norm can be represented either by a standalone norm module or by + a fused norm+linear projection. MLA can use the fused form; DSA cannot + because it needs the normalized Q/KV values outside the projection. + + Constraints: + - `qk_l2_norm` is unsupported for MLA/DSA. + - A standalone Q norm is only usable when `q_lora_rank` is set. + - Explicit norm modules cannot be paired with fused norm+linear projections. + - Disabled QK norm rejects both explicit norms and fused norm+linear projections. + - DSA with QK norm requires non-fused projections and standalone Q/KV norms. + """ + + def __init__(self, config: MLATransformerConfig, submodules) -> None: + """Capture the configuration, requested modules, and backend implementations.""" + self.config = config + self.submodules = submodules + self.has_q_lora = config.q_lora_rank is not None + self.is_dsa = config.experimental_attention_variant == "dsa" + self.variant_str = "DSA" if self.is_dsa else "MLA" + + backend = get_backend(config.transformer_impl) + self.qk_norm_impl = backend.layer_norm( + rms_norm=config.normalization == "RMSNorm", for_qk=True + ) + self.linear_impl = backend.column_parallel_linear() + self.fused_norm_linear_impl = backend.column_parallel_layer_norm_linear() + + def resolve(self) -> _QKNormResolvedConfig: + """Validate the specification and return the modules to instantiate. + + Returns: + The Q/KV norms and projections after applying the MLA or DSA constraints. + + Raises: + ValueError: If the requested norm placement is unsupported or conflicting. + """ + if self.config.qk_l2_norm: + raise ValueError(f"qk_l2_norm is not supported with {self.variant_str}.") + + self._reject_common_spec_conflicts() + if not self.config.qk_layernorm: + return self._resolve_disabled_qk_layernorm() + if self.is_dsa: + return self._resolve_dsa_qk_layernorm() + return self._resolve_mla_qk_layernorm() + + def _resolve_disabled_qk_layernorm(self) -> _QKNormResolvedConfig: + """Resolve projections when Q/KV normalization is disabled. + + Explicit norm modules and fused norm-linear projections are rejected because + they would still introduce Q/KV normalization. + """ + linear_q_proj_cls = IdentityOp + linear_q_up_proj_cls = IdentityOp + + if self.has_q_lora: + self._reject_disabled_norm( + self.submodules.linear_q_up_proj, + self.submodules.q_layernorm, + "linear_q_up_proj", + "q_layernorm", + ) + linear_q_up_proj_cls = self.submodules.linear_q_up_proj or self.linear_impl + else: + if self._is_fused_norm_linear(self.submodules.linear_q_proj): + raise ValueError( + f"spec sets linear_q_proj={self.submodules.linear_q_proj}, but " + "qk_layernorm/qk_l2_norm are supposed to be disabled" + ) + linear_q_proj_cls = self.submodules.linear_q_proj or self.linear_impl + + self._reject_disabled_norm( + self.submodules.linear_kv_up_proj, + self.submodules.kv_layernorm, + "linear_kv_up_proj", + "kv_layernorm", + ) + return self._result( + linear_q_proj=linear_q_proj_cls, + linear_q_up_proj=linear_q_up_proj_cls, + linear_kv_up_proj=self.submodules.linear_kv_up_proj or self.linear_impl, + q_layernorm=IdentityOp, + kv_layernorm=IdentityOp, + ) + + def _resolve_dsa_qk_layernorm(self) -> _QKNormResolvedConfig: + """Resolve DSA's standalone Q/KV norms and non-fused projections. + + DSA consumes the normalized Q/KV values outside the projection, so it cannot + use fused norm-linear projections. + """ + if not self.has_q_lora: + raise ValueError( + "`qk_layernorm=True` with `q_lora_rank is None` is not supported for DSA " + "because DSA cannot fuse Q norm into `linear_q_proj`." + ) + + return self._result( + linear_q_proj=IdentityOp, + linear_q_up_proj=self._dsa_linear_or_default( + self.submodules.linear_q_up_proj, "linear_q_up_proj" + ), + linear_kv_up_proj=self._dsa_linear_or_default( + self.submodules.linear_kv_up_proj, "linear_kv_up_proj" + ), + q_layernorm=self._default_if_trivial(self.submodules.q_layernorm, self.qk_norm_impl), + kv_layernorm=self._default_if_trivial(self.submodules.kv_layernorm, self.qk_norm_impl), + ) + + def _resolve_mla_qk_layernorm(self) -> _QKNormResolvedConfig: + """Resolve MLA norms, fusing them into projections when no norm is explicit.""" + q_norm_cls = self.submodules.q_layernorm or IdentityOp + linear_q_proj_cls = IdentityOp + linear_q_up_proj_cls = IdentityOp + + if self.has_q_lora: + if self._is_trivial(q_norm_cls): + linear_q_up_proj_cls = self._mla_fused_linear_or_default( + self.submodules.linear_q_up_proj, "linear_q_up_proj" + ) + else: + linear_q_up_proj_cls = self._non_fused_or_default( + self.submodules.linear_q_up_proj, "linear_q_up_proj" + ) + else: + linear_q_proj_cls = self._mla_fused_linear_or_default( + self.submodules.linear_q_proj, "linear_q_proj" + ) + + kv_norm_cls = self.submodules.kv_layernorm or IdentityOp + if self._is_trivial(kv_norm_cls): + linear_kv_up_proj_cls = self._mla_fused_linear_or_default( + self.submodules.linear_kv_up_proj, "linear_kv_up_proj" + ) + else: + linear_kv_up_proj_cls = self._non_fused_or_default( + self.submodules.linear_kv_up_proj, "linear_kv_up_proj" + ) + + return self._result( + linear_q_proj=linear_q_proj_cls, + linear_q_up_proj=linear_q_up_proj_cls, + linear_kv_up_proj=linear_kv_up_proj_cls, + q_layernorm=q_norm_cls, + kv_layernorm=kv_norm_cls, + ) + + def _reject_common_spec_conflicts(self) -> None: + """Reject conflicts that apply regardless of the selected attention variant.""" + if not self.has_q_lora and not self._is_trivial(self.submodules.q_layernorm): + self._raise_unused_q_norm() + if self.has_q_lora: + self._reject_explicit_norm_with_fused_linear( + self.submodules.linear_q_up_proj, + self.submodules.q_layernorm, + "linear_q_up_proj", + "q_layernorm", + ) + self._reject_explicit_norm_with_fused_linear( + self.submodules.linear_kv_up_proj, + self.submodules.kv_layernorm, + "linear_kv_up_proj", + "kv_layernorm", + ) + + def _reject_disabled_norm(self, module_spec, norm_spec, module_name, norm_name) -> None: + """Reject a norm module or fused projection when Q/KV norm is disabled.""" + if self._is_fused_norm_linear(module_spec) or not self._is_trivial(norm_spec): + raise ValueError( + f"spec sets {module_name}={module_spec} and " + f"{norm_name}={norm_spec}, but " + "qk_layernorm/qk_l2_norm are supposed to be disabled" + ) + + def _reject_explicit_norm_with_fused_linear( + self, module_spec, norm_spec, module_name, norm_name + ) -> None: + """Reject specifying the same norm both explicitly and inside a projection.""" + if not self._is_trivial(norm_spec) and self._is_fused_norm_linear(module_spec): + raise ValueError( + f"`{norm_name}={norm_spec}` is non-trivial " + f"and `{module_name}={module_spec}` is a " + f"fused norm+linear; either unset `{norm_name}` or use a " + f"linear layer without norm fusion for `{module_name}`" + ) + + def _non_fused_or_default(self, module_spec, module_name): + """Return a linear implementation, requiring it not to fuse normalization.""" + linear_cls = module_spec or self.linear_impl + self._require_linear(linear_cls, module_name) + if self._is_fused_norm_linear(linear_cls): + raise ValueError( + f"`{module_name}={module_spec}` is fused norm+linear, but a non-fused linear " + f"is required" + ) + return linear_cls + + def _dsa_linear_or_default(self, module_spec, module_name): + """Return DSA's non-fused projection implementation. + + This uses a DSA-specific diagnostic so the rejected constraint is clear. + """ + linear_cls = module_spec or self.linear_impl + self._require_linear(linear_cls, module_name) + if self._is_fused_norm_linear(linear_cls): + raise ValueError( + f"`{module_name}={module_spec}` is fused norm+linear, " + f"which is not supported for DSA." + ) + return linear_cls + + def _mla_fused_linear_or_default(self, module_spec, module_name): + """Return a fused MLA projection, using the backend default when available.""" + if self._is_fused_norm_linear(module_spec): + return module_spec + return self._require_linear(self.fused_norm_linear_impl, module_name) + + def _require_linear(self, module_spec, module_name): + """Return a configured projection or report that no viable implementation exists.""" + if module_spec is None: + raise RuntimeError( + "qk_layernorm requires TransformerEngine or " + "q_layernorm/kv_layernorm to be set in the spec " + f"to build `{module_name}`." + ) + return module_spec + + def _raise_unused_q_norm(self) -> NoReturn: + """Report an explicit Q norm that has no Q-LoRA projection to consume it.""" + help_msg = "" + if not self._is_fused_norm_linear(self.submodules.linear_q_proj): + help_msg = ( + f"Please use a fused norm+linear for " + f"`linear_q_proj={self.submodules.linear_q_proj}` if " + f"you intend to have a Q-norm." + ) + raise ValueError( + f"`q_layernorm={self.submodules.q_layernorm}` is non-trivial, " + f"but `q_lora_rank is None`, meaning it will not be used." + f"{help_msg}" + ) + + def _is_fused_norm_linear(self, module_spec) -> bool: + """Return whether a module specification selects the backend fused projection.""" + module_cls = module_spec.module if isinstance(module_spec, ModuleSpec) else module_spec + return self.fused_norm_linear_impl is not None and module_cls is self.fused_norm_linear_impl + + @staticmethod + def _is_trivial(module_spec) -> bool: + """Return whether a norm slot is unset or explicitly an identity operation.""" + return module_spec in (None, IdentityOp) + + @classmethod + def _default_if_trivial(cls, module_spec, default): + """Replace an unset or identity specification with the supplied default.""" + if cls._is_trivial(module_spec): + return default + return module_spec + + @staticmethod + def _result( + *, linear_q_proj, linear_q_up_proj, linear_kv_up_proj, q_layernorm, kv_layernorm + ) -> _QKNormResolvedConfig: + """Package the resolved Q/KV norms and projections in the caller's schema.""" + return dict( + linear_q_proj=linear_q_proj, + linear_q_up_proj=linear_q_up_proj, + linear_kv_up_proj=linear_kv_up_proj, + q_layernorm=q_layernorm, + kv_layernorm=kv_layernorm, + ) diff --git a/megatron/core/transformer/mlp.py b/megatron/core/transformer/mlp.py index 0a107a0b0cf..c165bf8e016 100644 --- a/megatron/core/transformer/mlp.py +++ b/megatron/core/transformer/mlp.py @@ -26,9 +26,15 @@ from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer.module import MegatronModule from megatron.core.transformer.transformer_config import TransformerConfig -from megatron.core.transformer.utils import cat_with_oom_fallback, sharded_state_dict_default +from megatron.core.transformer.utils import ( + cat_with_oom_fallback, + ensure_metadata_has_dp_cp_group, + sharded_state_dict_default, +) from megatron.core.typed_torch import apply_module, not_none from megatron.core.utils import ( + get_pg_rank, + get_pg_size, get_tensor_model_parallel_group_if_none, nvtx_range_pop, nvtx_range_push, @@ -356,6 +362,7 @@ def sharded_state_dict( self, prefix: str = "", sharded_offsets: tuple = (), metadata: Optional[dict] = None ) -> ShardedStateDict: """Return the sharded state dictionary of the module.""" + metadata = ensure_metadata_has_dp_cp_group(metadata) sharded_state_dict = {} singleton_local_shards = (metadata or {}).get('singleton_local_shards', False) for name, module in self._modules.items(): @@ -366,7 +373,11 @@ def sharded_state_dict( for k, v in sub_sd.items(): if k in (f"{prefix}{name}.weight", f"{prefix}{name}.bias"): sub_sd[k] = apply_swiglu_sharded_factory( - v, sharded_offsets, singleton_local_shards + v, + sharded_offsets, + singleton_local_shards, + tp_group=self.tp_group, + dp_group=metadata['dp_cp_group'], ) sharded_state_dict.update(sub_sd) return sharded_state_dict @@ -392,6 +403,14 @@ def as_mlp_submodule( assert hasattr( pg_collection, 'tp' ), 'TP process group is required for MLP in TransformerLayer' + + # fc1/fc2 resolve GTP_remat at the leaf; the TE op-fused MLP ignores shards, so fail fast. + if hasattr(cls, '_make_fused_impl'): + assert config.gtp_weight_remat_size <= 1, ( + f"{cls.__name__}: GTP sharding of the dense MLP is not supported with the " + "TE fused MLP / GroupedLinear path (_make_fused_impl ignores GTP shards). " + "Use the non-fused MLP submodule, or do not enable GTP for dense MLP layers." + ) return cls( config=config, submodules=submodules, @@ -405,7 +424,11 @@ def as_mlp_submodule( # pylint: disable=missing-function-docstring def apply_swiglu_sharded_factory( - original_sh_ten, sharded_offsets, singleton_local_shards: bool = False + original_sh_ten, + sharded_offsets, + singleton_local_shards: bool = False, + tp_group: torch.distributed.ProcessGroup | None = None, + dp_group: torch.distributed.ProcessGroup | None = None, ): # We must split the tensor into 2 parts, each sharded separately. # This requires a ShardedTensorFactory which `chunk`s during saving @@ -419,15 +442,24 @@ def apply_swiglu_sharded_factory( assert ( original_sh_ten.global_offset[swiglu_shard_axis + prepend_axis_num] % local_axis_size == 0 ) - rank_offset = ( - original_sh_ten.global_offset[swiglu_shard_axis + prepend_axis_num] // local_axis_size - ) axis_frag = original_sh_ten.axis_fragmentations[swiglu_shard_axis + prepend_axis_num] + # Only FSDP2 supports torch_dist ShardedTensor. (Add other DP sharding algos here if needed.) + is_torch_fsdp2_param = getattr(original_sh_ten, "is_torch_fsdp2_param", False) + if is_torch_fsdp2_param: + assert dp_group is not None + dp_size = get_pg_size(dp_group) + is_dp_sharded = dp_size > 1 + else: + is_dp_sharded = False + @torch.no_grad() def sh_ten_build_fn( key: str, t: torch.Tensor, replica_id: ReplicaId, flattened_range: Optional[slice] ): + rank_offset = ( + original_sh_ten.global_offset[swiglu_shard_axis + prepend_axis_num] // local_axis_size + ) if singleton_local_shards: offset_w = (swiglu_shard_axis + prepend_axis_num, rank_offset, axis_frag) offset_v = (swiglu_shard_axis + prepend_axis_num, rank_offset, axis_frag) @@ -463,10 +495,67 @@ def sh_ten_build_fn( ), ] + @torch.no_grad() + def dp_sh_ten_build_fn( + key: str, t: torch.Tensor, replica_id: ReplicaId, flattened_range: Optional[slice] + ): + assert not singleton_local_shards, ( + "FSDP does not support singleton ShardedTensor for SwiGLU fused FC1. " + "Set singleton_local_shards=False, which is the default in MCore." + ) + # FSDP shards the TP-local [W; V] SwiGLU FC1 tensor over DP along dim 0. + # TP sharding produces TP-sharded pairs of W/V, followed by DP sharding! + assert tp_group is not None and dp_group is not None + tp_size = get_pg_size(tp_group) + global_axis = swiglu_shard_axis + prepend_axis_num + tp_rank = get_pg_rank(tp_group) + dp_rank = get_pg_rank(dp_group) + # Size of a TP shard for W + V. + tp_local_axis_size = original_sh_ten.global_shape[global_axis] // tp_size + assert tp_local_axis_size % 2 == 0 # W and V should be symmetrically shaped. + # Size of a TP shard for W or V. "Half" size TP-shard. + half_axis_size = tp_local_axis_size // 2 + # Check that the TP-local W or V is cleanly divisible by DP. + assert half_axis_size % local_axis_size == 0, ( + "SwiGLU FC1 FSDP ShardedTensor requires each DP shard of " + "linear_fc1 to be completely inside either the W or V half." + ) + # Number of DP shards per W or V TP-shard, and make sure + # that DP sharding spans both W and V. + shards_per_half = half_axis_size // local_axis_size + assert dp_size == 2 * shards_per_half + + # Compute if DP rank maps to W or V in [W; V]. + swiglu_half_idx, half_dp_shard_idx = divmod(dp_rank, shards_per_half) + # If W, then 0. If V, then 1. + assert swiglu_half_idx in (0, 1) + # Map [ W; V ] to this rank's shard [ {W_tpx; V_tpx}_dpy ]. + shard_rank_offset = ( + # W or V half of the [W; V] global data. + swiglu_half_idx * tp_size * shards_per_half + # TP Shard Index + + tp_rank * shards_per_half + # TP-DP Shard Index + + half_dp_shard_idx + ) + + return [ + ShardedTensor.from_rank_offsets( + key, + t, + *sharded_offsets, + (global_axis, shard_rank_offset, axis_frag), + replica_id=replica_id, + prepend_axis_num=prepend_axis_num, + ) + ] + + # Construct a ShardedTensorFactory. + sh_ten_factory_build_function = dp_sh_ten_build_fn if is_dp_sharded else sh_ten_build_fn return ShardedTensorFactory( original_sh_ten.key, original_sh_ten.data, - sh_ten_build_fn, + sh_ten_factory_build_function, cat_with_oom_fallback, original_sh_ten.replica_id, flattened_range=original_sh_ten.flattened_range, diff --git a/megatron/core/transformer/module.py b/megatron/core/transformer/module.py index 35c5faab550..3b1a2b30e7e 100644 --- a/megatron/core/transformer/module.py +++ b/megatron/core/transformer/module.py @@ -197,12 +197,22 @@ def __init__(self, config: TransformerConfig, vp_stage: Optional[int] = None): self.cuda_graph_backward_dw_wrapper = None def init_backward_dw_wrapper(self): - """Initialize the backward_dw_wrapper.""" - from megatron.core.models.gpt.fine_grained_callables import _BackwardDWWrapper + """Initialize ``self.backward_dw_wrapper`` for delayed-wgrad scheduling. + + The wrapper coordinates the per-layer wgrad callables (attention + wgrad, optional shared-expert wgrad) with cuda-graph replay scope so + captured components are not re-run eagerly. The method is defined on + ``GraphableMegatronModule`` so any graphable subclass can opt in; + ``_BackwardDWWrapper`` itself currently asserts the underlying layer + is a ``TransformerLayer``, so MambaLayer-derived modules implement + ``backward_dw`` directly and skip this helper. + """ + from megatron.core.models.common.utils import _BackwardDWWrapper config = getattr(self, 'config', None) assert config is not None, ( - "TransformerLayer must be initialized before calling " "`init_backward_dw_wrapper`." + "Module must be fully constructed (config set) before calling " + "`init_backward_dw_wrapper`." ) self.backward_dw_wrapper = _BackwardDWWrapper(self) diff --git a/megatron/core/transformer/moe/experts.py b/megatron/core/transformer/moe/experts.py index ec9ba66e809..7a19ad6f900 100644 --- a/megatron/core/transformer/moe/experts.py +++ b/megatron/core/transformer/moe/experts.py @@ -61,9 +61,18 @@ import transformer_engine as te from megatron.core.extensions.transformer_engine import Fp8Padding, Fp8Unpadding + + try: + from transformer_engine.pytorch.ops.basic.grouped_linear import ( + GRAD_INPUT_BUFFER_KEY, + OUTPUT_BUFFER_KEY, + ) + except ImportError: + GRAD_INPUT_BUFFER_KEY = OUTPUT_BUFFER_KEY = None else: te = None # type: ignore[assignment, misc] Fp8Padding, Fp8Unpadding = None, None + GRAD_INPUT_BUFFER_KEY = OUTPUT_BUFFER_KEY = None try: import flashinfer.fused_moe as fused_moe @@ -363,6 +372,12 @@ def _activation_name(): return _unsupported(f"linear_fc1 is {type(self.linear_fc1).__name__}") if not isinstance(self.linear_fc2, te.pytorch.GroupedLinear): return _unsupported(f"linear_fc2 is {type(self.linear_fc2).__name__}") + if ( + self.linear_fc2.use_bias + and "scale_bias" not in inspect.signature(te_ops.GroupedLinear.__init__).parameters + ): + # Older TE op-fuser versions cannot scale FC2 bias by router probabilities + return _unsupported("TE op-fuser too old to scale FC2 bias by router probabilities") # Check activation: SwiGLU, quick GEGLU, or weighted squared ReLU. # Clamped SwiGLU (e.g. DSv4) routes through ScaledClampedQGeGLU with @@ -396,6 +411,8 @@ def _activation_name(): elif self.config.activation_func == squared_relu: if not hasattr(te_ops, "ScaledSReLU"): return _unsupported("weighted squared_relu needs ScaledSReLU") + else: + return _unsupported(f"unsupported activation {_activation_name()}") # Check TE CuTe DSL fused kernel conditions (must match TE's # fuse_grouped_mlp_ops matching logic). @@ -568,6 +585,7 @@ def register_grouped_linear_params( ops.append(op) # FC2 + fc2_bias_kwargs = {"scale_bias": True} if self.linear_fc2.use_bias else {} op = te.pytorch.ops.GroupedLinear( self.linear_fc2.num_gemms, self.linear_fc2.in_features, @@ -579,6 +597,8 @@ def register_grouped_linear_params( single_grouped_weight=fc2_single_grouped_weight, single_grouped_bias=fc2_single_grouped_bias, delay_wgrad_compute=fc2_delay_wgrad_compute, + # Preserve p * (FC2(x) + bias) after the scaled activation moves p before FC2. + **fc2_bias_kwargs, ) # In single grouped mode, clear stale per-expert meta params so TE does not reset @@ -598,7 +618,7 @@ def _make_fused_impl_pre_forward_hook(self) -> Callable: """Make function that calls submodule pre-forward callback hooks. This is intended for compatibility with - DistributedDataParallel hooks that trigger parameter + DistributedDataParallel/FSDP hooks that trigger parameter all-gathers. It does not support general pre-forward hooks since they may manipulate intermediate tensors that are never instantiated by the fused implementation. @@ -616,6 +636,7 @@ def forward_pre_hook(module, *_) -> None: f"but a {submodule.__class__.__name__} submodule " "has a pre-forward hook that modifies the input tensor." ) + self._ensure_main_grad_for_fused_impl() return forward_pre_hook @@ -640,11 +661,30 @@ def forward_post_hook(_module, _inputs, output): return forward_post_hook + @staticmethod + def _ensure_main_grad(linear_module: torch.nn.Module) -> None: + """Expose FSDP main_grad buffers required by TE fused wgrad accumulation.""" + if not getattr(linear_module, "fuse_wgrad_accumulation", False): + return + for param in linear_module.parameters(recurse=False): + get_main_grad = getattr(param, "get_main_grad", None) + if get_main_grad is not None and getattr(param, "main_grad", None) is None: + param.main_grad = get_main_grad() + if hasattr(param, "overwrite_main_grad"): + param.overwrite_main_grad = True + + def _ensure_main_grad_for_fused_impl(self) -> None: + """Expose wrapper parameter main_grad buffers before TE fused ops run.""" + self._ensure_main_grad(self.linear_fc1) + self._ensure_main_grad(self.linear_fc2) + def _fused_forward( self, permuted_local_hidden_states: torch.Tensor, tokens_per_expert: torch.Tensor, permuted_probs: torch.Tensor, + output_buffer: Optional[torch.Tensor] = None, + grad_input_buffer: Optional[torch.Tensor] = None, ) -> torch.Tensor: """Forward pass using Transformer Engine operation fuser API.""" @@ -701,16 +741,35 @@ def _fused_forward( fine_grained_activation_offloading, permuted_local_hidden_states, offload_name ) with fused_group_mlp_manager as permuted_local_hidden_states: + # NCCL-EP zero-copy is active exactly when ``output_buffer`` is not None, and then the + # fused-MLP input aliases the persistent symm buffer (also the fc2 output combine + # reads), whose storage is non-resizable — so skip the force-release in that case. forced_released_tensors = ( - [permuted_local_hidden_states] if fine_grained_activation_offloading else [] + [permuted_local_hidden_states] + if fine_grained_activation_offloading and output_buffer is None + else [] ) with stash_context: + # NCCL-EP zero-copy: route the fc2 output (fwd combine reads it one-sided) and the + # fc1 dgrad (bwd dispatch scatters it one-sided) into caller-provided symm buffers. + # op_kwargs keys are basic-op indices into [fc1, activation, fc2]: 0=fc1, -1=fc2. + op_kwargs = {} + if output_buffer is not None: + op_kwargs[-1] = {OUTPUT_BUFFER_KEY: output_buffer} + if grad_input_buffer is not None: + op_kwargs[0] = {GRAD_INPUT_BUFFER_KEY: grad_input_buffer} # Call fused impl + fc2_extra_inputs = ( + (tokens_per_expert, permuted_probs) + if self.linear_fc2.use_bias + else (tokens_per_expert,) + ) output = ops( permuted_local_hidden_states, tokens_per_expert, # FC1 permuted_probs, # Scaled activation - tokens_per_expert, # FC2 + *fc2_extra_inputs, # FC2 splits and, for bias, its per-token scale + **({"op_kwargs": op_kwargs} if op_kwargs else {}), ) output = fused_group_mlp_manager.group_offload( output, @@ -738,6 +797,8 @@ def forward( permuted_local_hidden_states: torch.Tensor, tokens_per_expert: torch.Tensor, permuted_probs: torch.Tensor, + output_buffer: Optional[torch.Tensor] = None, + grad_input_buffer: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, Optional[torch.Tensor]]: """Forward of TEGroupedMLP @@ -746,6 +807,10 @@ def forward( local experts. tokens_per_expert (torch.Tensor): The number of tokens per expert. permuted_probs (torch.Tensor): The permuted probs of each token produced by the router. + output_buffer (torch.Tensor, optional): Preallocated buffer to write the fc2 output into + (NCCL-EP zero-copy fwd combine); only the fused op-fuser path supports it. + grad_input_buffer (torch.Tensor, optional): Preallocated buffer to write the fc1 dgrad + into (NCCL-EP zero-copy bwd dispatch); only the fused op-fuser path supports it. Return: output (torch.Tensor): The output of the local experts. @@ -754,10 +819,17 @@ def forward( # Call fused impl if enabled if self._with_fused_impl: output = self._fused_forward( - permuted_local_hidden_states, tokens_per_expert, permuted_probs + permuted_local_hidden_states, + tokens_per_expert, + permuted_probs, + output_buffer, + grad_input_buffer, ) output_bias = None return output, output_bias + assert ( + output_buffer is None and grad_input_buffer is None + ), "output_buffer/grad_input_buffer require the TE op-fuser (fused) path" # Apply padding if needed unpadded_tokens_per_expert = None @@ -937,7 +1009,11 @@ def sharded_state_dict( for k in (f'{name}.weight{i}', f'{name}.bias{i}'): if k in sub_sd: sub_sd[k] = apply_swiglu_sharded_factory( - sub_sd[k], new_sharded_offsets, singleton_local_shards + sub_sd[k], + new_sharded_offsets, + singleton_local_shards, + tp_group=self.tp_group, + dp_group=metadata['dp_cp_group'], ) if singleton_local_shards: replace_prefix_for_sharding(sub_sd, '', f'{prefix}experts.') diff --git a/megatron/core/transformer/moe/fused_a2a.py b/megatron/core/transformer/moe/fused_a2a.py index 2a6810e0848..4285ae36298 100644 --- a/megatron/core/transformer/moe/fused_a2a.py +++ b/megatron/core/transformer/moe/fused_a2a.py @@ -892,6 +892,12 @@ def nccl_ep_finalize(): if HAVE_TE_EP: + def alloc_ep_symm_buffer(shape, dtype, ep_group): + """Allocate one persistent NCCL symm-mem buffer (per-buffer collective rendezvous). mcore's + zero-copy buffers are all persistent and non-pool; the symm mem-pool is used only by TE for + the per-call recv buffers it recycles.""" + return te_ep.symm_mem_alloc(shape, dtype, ep_group) + def new_nccl_ep_buffer( top_k, max_tokens_per_rank, @@ -902,8 +908,9 @@ def new_nccl_ep_buffer( ): """Build a fresh TE EpBuffer for one dispatch/combine pair. - The buffer owns handle_mem (the routing table dispatch writes and combine reads) and - the receive buffers; a new one is built per dispatch and dropped after combine. + The buffer owns handle_mem (the routing table dispatch writes and combine reads); a new one + is built per dispatch and dropped after combine. Payload symm buffers are not owned here — + they are caller-supplied to dispatch/combine or allocated on the fly by TE. """ return te_ep.EpBuffer( top_k=top_k, @@ -914,7 +921,9 @@ def new_nccl_ep_buffer( alignment=alignment, ) - def nccl_ep_dispatch(buffer, tokens, topk_idx, topk_weights): + def nccl_ep_dispatch( + buffer, tokens, topk_idx, topk_weights, recv_tokens=None, recv_topk_weights=None + ): """Autograd-aware prepare + dispatch via TransformerEngine NCCL EP. Args: @@ -924,6 +933,9 @@ def nccl_ep_dispatch(buffer, tokens, topk_idx, topk_weights): topk_idx (torch.Tensor): ``int64`` ``[num_local_tokens, top_k]`` global expert ids per token. topk_weights (torch.Tensor): ``float32`` ``[num_local_tokens, top_k]`` weights. + recv_tokens, recv_topk_weights (torch.Tensor, optional): caller-owned symm dispatch + recv buffers (fp8 zero-copy). Left None, TE allocates them (bf16 zero-copy: symm + mem-pool; normal: plain). Returns: tuple: ``(recv_tokens, tokens_per_expert, dispatched_probs)``: @@ -938,11 +950,16 @@ def nccl_ep_dispatch(buffer, tokens, topk_idx, topk_weights): ``tokens_per_expert`` is non-differentiable. """ recv_tokens, dispatched_probs, tokens_per_expert = te_ep.ep_dispatch( - buffer, tokens, topk_idx, topk_weights + buffer, + tokens, + topk_idx, + topk_weights, + recv_tokens=recv_tokens, + recv_topk_weights=recv_topk_weights, ) return recv_tokens, tokens_per_expert, dispatched_probs - def nccl_ep_combine(buffer, expert_out, num_local_tokens=None): + def nccl_ep_combine(buffer, expert_out, num_local_tokens=None, grad_out=None): """Autograd-aware combine via TransformerEngine NCCL EP (no scatter step). Args: @@ -951,14 +968,20 @@ def nccl_ep_combine(buffer, expert_out, num_local_tokens=None): already weighted. num_local_tokens (int): Rows of the result (local token count for this forward). When None, TE uses ``buffer.max_tokens_per_rank``. + grad_out (torch.Tensor, optional): caller-owned symm buffer the backward scatters the + expert_out grad into (zero-copy). Left None, TE allocates it (bf16: symm mem-pool; + normal: plain). Returns: torch.Tensor: ``[num_local_tokens, hidden]`` combined output, in local token order. """ - return te_ep.ep_combine(buffer, expert_out, num_local_tokens=num_local_tokens) + return te_ep.ep_combine( + buffer, expert_out, num_local_tokens=num_local_tokens, grad_out=grad_out + ) else: + alloc_ep_symm_buffer = None new_nccl_ep_buffer = None nccl_ep_dispatch = None nccl_ep_combine = None diff --git a/megatron/core/transformer/moe/inference_routing_mask_kernel.py b/megatron/core/transformer/moe/inference_routing_mask_kernel.py new file mode 100644 index 00000000000..e38869a1f6d --- /dev/null +++ b/megatron/core/transformer/moe/inference_routing_mask_kernel.py @@ -0,0 +1,101 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Triton kernel for masking CUDA-graph padding rows of a local routing map. + +Under CUDA-graph capture the local token count is padded up to a captured +graph size; those padding rows have garbage routing indices and, if left +alone, would dispatch padding tokens to real experts. This kernel zeroes +that out by writing ``-1`` into every topk slot of rows in +``[real_token_count, local_tokens)``. + +The kernel reads ``real_token_count`` from a fixed-address ``int32[1]`` GPU +tensor, so it is safe to call from inside a captured graph: only the value +behind the pointer changes between replays. +""" + +from torch import Tensor + +try: + import triton + import triton.language as tl + + HAVE_TRITON = True +except ImportError: + from unittest.mock import MagicMock + + from megatron.core.utils import null_decorator + + triton = MagicMock() + triton.jit = null_decorator + triton.autotune = null_decorator + tl = MagicMock() + HAVE_TRITON = False + + +@triton.jit +def _mask_routing_padding_kernel( + routing_map_ptr, # int64* [total_rows, topk] + real_token_count_ptr, # int32* [1] + total_rows: tl.int32, + tp_rank: tl.int32, # SP/TP rank — local row r maps to global row r + tp_rank*total_rows + TOPK: tl.constexpr, # actual topk + BLOCK_M: tl.constexpr, # rows per program + BLOCK_TOPK: tl.constexpr, # next_power_of_2(TOPK), column block +): + """Fill `routing_map[real_token_count:, :]` with -1, BLOCK_M rows per program.""" + pid = tl.program_id(0) + rows = pid * BLOCK_M + tl.arange(0, BLOCK_M) + + real_count = tl.load(real_token_count_ptr).to(tl.int32) + + # real_count is in the global (pre-SP-shard) frame; rows is local to this SP rank. + global_rows = rows + tp_rank * total_rows + row_mask = (global_rows >= real_count) & (rows < total_rows) + + cols = tl.arange(0, BLOCK_TOPK) + col_mask = cols < TOPK + + offs = rows[:, None].to(tl.int64) * TOPK + cols[None, :].to(tl.int64) + mask = row_mask[:, None] & col_mask[None, :] + + neg_one = tl.full((BLOCK_M, BLOCK_TOPK), -1, dtype=tl.int64) + tl.store(routing_map_ptr + offs, neg_one, mask=mask) + + +def mask_routing_padding( + routing_map: Tensor, real_token_count_tensor: Tensor, tp_rank: int = 0 +) -> None: + """In-place fill -1 into ``routing_map[real_token_count:, :]``. + + Args: + routing_map: ``[N, topk]`` int64 local routing map. ``N`` is the + (possibly CUDA-graph-padded) local token count. + real_token_count_tensor: ``[1]`` int32 GPU tensor holding the real + (unpadded) token count for this step, in the global (pre-SP-shard) + frame. Read inside the kernel so the mask boundary moves correctly + across CUDA-graph replays. + tp_rank: This rank's index in the SP/TP group. Local row ``r`` is + row ``r + tp_rank * N`` in the global frame; the kernel uses this + offset to compare against ``real_token_count_tensor``. + """ + assert routing_map.is_cuda, "routing_map must be on CUDA" + assert routing_map.dim() == 2, f"expected 2D routing_map, got {routing_map.shape}" + assert routing_map.dtype.is_floating_point is False, "routing_map must be integer" + + total_rows, topk = routing_map.shape + if total_rows == 0: + return + + BLOCK_M = 8 if total_rows < 64 else 128 + BLOCK_TOPK = triton.next_power_of_2(topk) + grid = (triton.cdiv(total_rows, BLOCK_M),) + + _mask_routing_padding_kernel[grid]( + routing_map, + real_token_count_tensor, + total_rows=total_rows, + tp_rank=tp_rank, + TOPK=topk, + BLOCK_M=BLOCK_M, + BLOCK_TOPK=BLOCK_TOPK, + ) diff --git a/megatron/core/transformer/moe/moe_layer.py b/megatron/core/transformer/moe/moe_layer.py index 59684a34b0d..880788139fb 100644 --- a/megatron/core/transformer/moe/moe_layer.py +++ b/megatron/core/transformer/moe/moe_layer.py @@ -616,8 +616,17 @@ def routed_experts_compute(self, hidden_states: torch.Tensor, probs: torch.Tenso dispatched_input, tokens_per_expert, permuted_probs, routing_map=routing_map ) else: + # NCCL-EP zero-copy: experts write fc2 output and fc1 dgrad straight into the combine / + # dispatch symm buffers. Passed only when set (non-TEGroupedMLP experts don't accept + # these kwargs). + output_buffer, grad_input_buffer = self.token_dispatcher.get_expert_zero_copy_buffers() + expert_kwargs = {} + if output_buffer is not None: + expert_kwargs["output_buffer"] = output_buffer + if grad_input_buffer is not None: + expert_kwargs["grad_input_buffer"] = grad_input_buffer expert_output, mlp_bias = apply_module(self.experts)( - dispatched_input, tokens_per_expert, permuted_probs + dispatched_input, tokens_per_expert, permuted_probs, **expert_kwargs ) assert mlp_bias is None, f"mlp_bias is not supported for {type(self.token_dispatcher)}" output = self.token_dispatcher.combine_preprocess(expert_output) diff --git a/megatron/core/transformer/moe/moe_logging.py b/megatron/core/transformer/moe/moe_logging.py index 16b60f66276..bb51688396d 100644 --- a/megatron/core/transformer/moe/moe_logging.py +++ b/megatron/core/transformer/moe/moe_logging.py @@ -621,12 +621,17 @@ def _sync_metrics( """ if pg_collection is None: pp_group = parallel_state.get_pipeline_model_parallel_group() + dp_group = None + else: + pp_group = pg_collection.pp + dp_group = getattr(pg_collection, 'dp_cp_gtp_remat', None) + # The metric DP-average must span gtp_remat peers (they hold distinct tokens), else the + # displayed value is a 1/gtp_remat subsample and looks noisy. Use the gtp_remat-inclusive + # group; CP ranks (already summed in reduce_group) average as a no-op. + if dp_group is None: dp_group = parallel_state.get_data_parallel_group( with_context_parallel=False, partial_data_parallel=False ) - else: - pp_group = pg_collection.pp - dp_group = pg_collection.dp for name in metric_names: if name not in self._metrics: diff --git a/megatron/core/transformer/moe/moe_utils.py b/megatron/core/transformer/moe/moe_utils.py index c8e197d2a3b..cfe44d74751 100644 --- a/megatron/core/transformer/moe/moe_utils.py +++ b/megatron/core/transformer/moe/moe_utils.py @@ -202,6 +202,44 @@ def sinkhorn(cost: torch.Tensor, tol: float = 0.0001) -> torch.Tensor: return d1 * cost * d0.unsqueeze(1) +def qb_dual_update( + scores: torch.Tensor, k: int, beta: torch.Tensor, update_beta: bool = True +) -> Tuple[torch.Tensor, torch.Tensor]: + """Dual coordinate-descent quantile-balancing routing assignment. + + Picks the top-k experts per token from ``scores - beta``. When ``update_beta`` is + True, also returns the raw column quantile of ``scores`` that drives each expert + toward ``m * k / n`` tokens. + + Args: + scores (torch.Tensor): Scores of shape ``[m, n]`` (tokens, experts). + k (int): Experts to select per token. + beta (torch.Tensor): Current per-expert bias of shape ``[n]``. + update_beta (bool): If False, return ``beta`` unchanged (eval/inference). + + Returns: + Tuple[torch.Tensor, torch.Tensor]: indices of shape ``[m, k]`` and either + ``beta`` (when ``update_beta`` is False) or the column quantile ``[n]``. + """ + num_tokens, num_experts = scores.shape + + topk_result = (scores - beta).topk(k + 1, dim=1) + indices = topk_result.indices[:, :-1] + + if not update_beta: + return indices, beta + + assert (num_tokens * k) % num_experts == 0, ( + "Quantile balancing requires the number of routed assignments " + f"({num_tokens} tokens * top-{k}) to be divisible by " + f"{num_experts} experts." + ) + col_target = num_tokens * k // num_experts + alpha = topk_result.values[:, -1:] + beta_local = (scores - alpha).topk(col_target + 1, dim=0).values[-1].contiguous() + return indices, beta_local + + def get_capacity( num_tokens: int, num_experts: int, capacity_factor: float, min_capacity: Optional[int] = None ) -> int: @@ -692,6 +730,7 @@ def topk_routing_with_score_function( fused: bool = False, router_replay: Optional['RouterReplay'] = None, dense_output: bool = False, + precomputed_indices: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: """Compute the routing probabilities and map for top-k selection with score function. @@ -716,6 +755,10 @@ def topk_routing_with_score_function( Defaults to None. dense_output (bool, optional): If True, return dense tensors [num_tokens, topk] instead of sparse tensors [num_tokens, num_experts]. Defaults to False. + precomputed_indices (torch.Tensor, optional): Top-k indices [num_tokens, topk] + selected by the caller. When given, the score function's + own top-k is bypassed and probs are computed at these + indices (e.g. for quantile balancing). Defaults to None. Returns: Tuple[torch.Tensor, torch.Tensor]: @@ -734,6 +777,9 @@ def topk_routing_with_score_function( """ assert logits.dim() == 2, f"Expected 2D logits [num_tokens, num_experts], got {logits.dim()}." num_tokens, num_experts = logits.shape + assert not ( + fused and precomputed_indices is not None + ), "precomputed_indices is not supported with the fused top-k score function." if fused: if not HAVE_TE or fused_topk_with_score_function is None: raise ValueError( @@ -804,16 +850,27 @@ def compute_topk(scores, topk, num_groups=None, group_topk=None): if score_function == "softmax": if use_pre_softmax: scores = torch.softmax(logits, dim=-1, dtype=torch.float32) - probs, top_indices = compute_topk(scores, topk, num_groups, group_topk) + if precomputed_indices is not None: + top_indices = precomputed_indices + probs = torch.gather(scores, dim=1, index=top_indices) + else: + probs, top_indices = compute_topk(scores, topk, num_groups, group_topk) else: - scores, top_indices = compute_topk(logits, topk, num_groups, group_topk) + if precomputed_indices is not None: + top_indices = precomputed_indices + scores = torch.gather(logits, dim=1, index=top_indices) + else: + scores, top_indices = compute_topk(logits, topk, num_groups, group_topk) probs = torch.softmax(scores, dim=-1, dtype=torch.float32) elif score_function in ("sigmoid", "sqrtsoftplus"): if score_function == "sigmoid": scores = torch.sigmoid(logits.float()) else: scores = torch.nn.functional.softplus(logits.float()).sqrt() - if expert_bias is not None: + if precomputed_indices is not None: + top_indices = precomputed_indices + scores = torch.gather(scores, dim=1, index=top_indices) + elif expert_bias is not None: scores_for_routing = scores + expert_bias.float() _, top_indices = compute_topk(scores_for_routing, topk, num_groups, group_topk) scores = torch.gather(scores, dim=1, index=top_indices) @@ -1473,7 +1530,10 @@ def get_default_pg_collection() -> ProcessGroupCollection: pg_collection.tp = parallel_state.get_tensor_model_parallel_group() pg_collection.cp = parallel_state.get_context_parallel_group() pg_collection.expt_tp = parallel_state.get_expert_tensor_parallel_group() - pg_collection.expt_dp = parallel_state.get_expert_data_parallel_group() + pg_collection.expt_dp = parallel_state.get_expert_data_parallel_group(with_gtp_remat=False) + pg_collection.expt_dp_gtp_remat = parallel_state.get_expert_data_parallel_group( + check_initialized=False + ) pg_collection.tp_ep = parallel_state.get_expert_tensor_and_model_parallel_group() pg_collection.tp_cp = parallel_state.get_tensor_and_context_parallel_group() pg_collection.tp_dp_cp = parallel_state.get_tensor_and_data_parallel_group( diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index 2796bc676a7..43e05287e4d 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -19,6 +19,7 @@ apply_router_token_dropping, compute_routing_scores_for_aux_loss, get_tokens_per_expert_and_token_count, + qb_dual_update, router_gating_linear, sinkhorn, switch_load_balancing_loss_func, @@ -263,6 +264,41 @@ def __init__( self.global_tokens_per_expert = None self.ga_steps = None + # Quantile balancing replaces the aux loss with a per-expert bias `qb_beta`. + # `qb_beta_accum`/`qb_beta_count` collect the per-microbatch quantile, reduced + # and reset each global batch. + if self.routing_type == "quantile_balancing": + assert not self.is_aux_loss_enabled(), ( + "Quantile balancing handles load balance via the bias update; " + "aux losses must be disabled (set moe_aux_loss_coeff to 0)." + ) + self.register_buffer( + 'qb_beta', + torch.zeros( + self.config.num_moe_experts, + dtype=torch.float32, + device=torch.cuda.current_device(), + ), + ) + self.register_buffer( + 'qb_beta_accum', + torch.zeros( + self.config.num_moe_experts, + dtype=torch.float32, + device=torch.cuda.current_device(), + ), + persistent=False, + ) + self.register_buffer( + 'qb_beta_count', + torch.zeros((), dtype=torch.long, device=torch.cuda.current_device()), + persistent=False, + ) + else: + self.qb_beta = None + self.qb_beta_accum = None + self.qb_beta_count = None + self.router_replay = None if self.config.moe_enable_routing_replay: self.router_replay = RouterReplay() @@ -277,6 +313,13 @@ def _maintain_float32_expert_bias(self): if hasattr(self, 'expert_bias') and self.expert_bias is not None: if self.expert_bias.dtype != torch.float32: self.expert_bias.data = self.expert_bias.data.to(torch.float32) + # Keep the QB bias in fp32 for the same reason. + if hasattr(self, 'qb_beta') and self.qb_beta is not None: + if self.qb_beta.dtype != torch.float32: + self.qb_beta.data = self.qb_beta.data.to(torch.float32) + if hasattr(self, 'qb_beta_accum') and self.qb_beta_accum is not None: + if self.qb_beta_accum.dtype != torch.float32: + self.qb_beta_accum.data = self.qb_beta_accum.data.to(torch.float32) def sinkhorn_load_balancing(self, logits: torch.Tensor): """Apply sinkhorn routing to the logits tensor. @@ -311,6 +354,81 @@ def _sinkhorn_activation(logits): scores = logits * map return scores, map + def quantile_balancing(self, logits: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """Apply quantile-balancing (QB) routing to the logits tensor. + + Selects top-k experts per token using a dual coordinate-descent update on + a per-expert bias ``qb_beta``. Load balance is handled entirely by the bias + update; auxiliary losses must be disabled when QB is active. + + Args: + logits (torch.Tensor): The logits tensor, shape ``[num_tokens, num_experts]``. + + Returns: + Tuple[torch.Tensor, torch.Tensor]: Sparse routing probs and boolean + routing map, each shaped ``[num_tokens, num_experts]``. + """ + assert ( + not self.config.moe_router_fusion + ), "Quantile balancing routing does not support moe_router_fusion." + assert ( + self.config.moe_router_num_groups is None and self.config.moe_router_group_topk is None + ), "Quantile balancing routing does not support group-limited routing." + + local_num_tokens = logits.shape[0] + # Gather logits across TP/CP so the quantile sees a whole sequence's tokens. + # The DP reduction and qb_beta update run at the global-batch boundary in + # finalize_model_grads._update_router_qb_beta. + gather_group = self.tp_cp_group + gather_size = gather_group.size() if gather_group is not None else 1 + + should_update_beta = self.training and torch.is_grad_enabled() + + with torch.no_grad(): + logits_fp32 = logits.detach().to(dtype=torch.float32) + + if gather_size > 1: + full_logits = torch.empty( + (local_num_tokens * gather_size, self.config.num_moe_experts), + dtype=logits_fp32.dtype, + device=logits_fp32.device, + ) + torch.distributed.all_gather_into_tensor( + full_logits, logits_fp32.contiguous(), group=gather_group + ) + gather_rank = torch.distributed.get_rank(group=gather_group) + else: + full_logits = logits_fp32 + gather_rank = 0 + + # Route with the previous batch's qb_beta; in training, accumulate this + # microbatch's quantile for the next update. + full_indices, beta_local = qb_dual_update( + full_logits, self.topk, self.qb_beta, update_beta=should_update_beta + ) + if should_update_beta: + self.qb_beta_accum.add_(beta_local) + self.qb_beta_count.add_(1) + + # Take this rank's rows (all_gather orders rows by rank). + if gather_size > 1: + indices = full_indices[ + gather_rank * local_num_tokens : (gather_rank + 1) * local_num_tokens + ].contiguous() + else: + indices = full_indices + + # QB only picks the experts; reuse the shared score function for the probs. + return topk_routing_with_score_function( + logits, + self.topk, + use_pre_softmax=self.config.moe_router_pre_softmax, + scaling_factor=self.config.moe_router_topk_scaling_factor, + score_function=self.score_function, + fused=self.config.moe_router_fusion, + precomputed_indices=indices, + ) + def get_aux_loss_coeff(self, aux_loss_type: str) -> float: """Return the aux loss coeff for the given auxiliary loss type. If the auxiliary loss type is not found, return 0.0. @@ -683,7 +801,11 @@ def apply_z_loss(self, logits, padding_mask: Optional[torch.Tensor] = None): layer_number = self.layer_number get_moe_metrics_tracker().record( - "z_loss", z_loss_mean / mtp_loss_scale, layer_number, num_layers + "z_loss", + z_loss_mean / mtp_loss_scale, + layer_number, + num_layers, + avg_group=self.tp_dp_cp_group, ) return logits @@ -815,6 +937,11 @@ def routing( probs, routing_map = self._hash_routing(logits, input_ids) elif self.routing_type == "sinkhorn": probs, routing_map = self.sinkhorn_load_balancing(logits) + elif self.routing_type == "quantile_balancing": + assert ( + padding_mask is None + ), "Quantile balancing routing does not support padding masks yet." + probs, routing_map = self.quantile_balancing(logits) else: probs, routing_map = topk_routing_with_score_function( logits, @@ -1008,6 +1135,7 @@ def _compiled_topk_routing( fused, router_replay, dense_output, + precomputed_indices, ): return topk_routing_with_score_function( logits, @@ -1021,11 +1149,17 @@ def _compiled_topk_routing( fused=fused, router_replay=router_replay, dense_output=dense_output, + precomputed_indices=precomputed_indices, ) def _forward(self, input: torch.Tensor, padding_mask: Optional[torch.Tensor] = None): logits = self.gating(input).squeeze(1) # [num_tokens, num_experts] + # QB selects on (logits - qb_beta); at inference qb_beta is fixed, so it's per-token. + precomputed_indices = None + if self.qb_beta is not None: + precomputed_indices = (logits - self.qb_beta).topk(self.topk, dim=1).indices + probs, top_indices = self._compiled_topk_routing( logits, self.topk, @@ -1038,6 +1172,7 @@ def _forward(self, input: torch.Tensor, padding_mask: Optional[torch.Tensor] = N fused=self.config.moe_router_fusion, router_replay=self.router_replay, dense_output=True, + precomputed_indices=precomputed_indices, ) return probs.squeeze(1), top_indices.squeeze(1) diff --git a/megatron/core/transformer/moe/shared_experts.py b/megatron/core/transformer/moe/shared_experts.py index 50e2ef6c0ce..06ac88b9fb7 100644 --- a/megatron/core/transformer/moe/shared_experts.py +++ b/megatron/core/transformer/moe/shared_experts.py @@ -398,6 +398,8 @@ def __init__( ) self._fused_grouped_swiglu_ops = None self._fused_grouped_swiglu_recipe = None + self._fused_grouped_swiglu_unit_scale = None + self._fused_grouped_swiglu_tokens_per_expert = {} self._validate_fused_grouped_swiglu() def _validate_fused_grouped_swiglu(self) -> None: @@ -483,7 +485,12 @@ def _make_fused_grouped_swiglu_ops(self) -> torch.nn.Module: op._glu_interleave_size = glu_interleave_size ops.append(op) - ops.append(te.pytorch.ops.ScaledSwiGLU(glu_interleave_size=glu_interleave_size)) + activation_op = te.pytorch.ops.ScaledSwiGLU(glu_interleave_size=glu_interleave_size) + # Shared experts are not router-gated. Mark this fused-op instance so + # TE can omit the optional forward cuDNN probability tensor without + # changing the semantics of routed single-group MLPs. + activation_op._grouped_mlp_unit_activation_scale = True + ops.append(activation_op) fc2_weight = self.linear_fc2.weight op = te.pytorch.ops.GroupedLinear( @@ -521,10 +528,21 @@ def _fused_grouped_swiglu_no_comm(self, hidden_states: torch.Tensor) -> torch.Te hidden_size = hidden_states.size(-1) hidden_states_2d = hidden_states.view(-1, hidden_size) total_tokens = hidden_states_2d.size(0) - tokens_per_expert = torch.full( - (1,), total_tokens, dtype=torch.long, device=hidden_states.device - ) - scales = torch.ones(total_tokens, device=hidden_states.device, dtype=hidden_states.dtype) + tokens_key = (hidden_states.device, total_tokens) + tokens_per_expert = self._fused_grouped_swiglu_tokens_per_expert.get(tokens_key) + if tokens_per_expert is None: + tokens_per_expert = torch.tensor( + [total_tokens], dtype=torch.long, device=hidden_states.device + ) + self._fused_grouped_swiglu_tokens_per_expert[tokens_key] = tokens_per_expert + scales = self._fused_grouped_swiglu_unit_scale + if ( + scales is None + or scales.device != hidden_states.device + or scales.dtype != hidden_states.dtype + ): + scales = torch.ones(1, device=hidden_states.device, dtype=hidden_states.dtype) + self._fused_grouped_swiglu_unit_scale = scales recipe = self._get_fused_grouped_swiglu_recipe() if self._fused_grouped_swiglu_ops is None: diff --git a/megatron/core/transformer/moe/token_dispatcher.py b/megatron/core/transformer/moe/token_dispatcher.py index 523f8af3d4d..6a48334cd94 100644 --- a/megatron/core/transformer/moe/token_dispatcher.py +++ b/megatron/core/transformer/moe/token_dispatcher.py @@ -2,6 +2,7 @@ import logging import os +import warnings from abc import ABC, abstractmethod from typing import List, Optional, Tuple @@ -20,6 +21,7 @@ from megatron.core.transformer.enums import CudaGraphModule from megatron.core.transformer.moe.fused_a2a import ( HYBRIDEP_TOKEN_ALIGNMENT, + alloc_ep_symm_buffer, deepepv2_combine, deepepv2_dispatch, ensure_nccl_ep_bootstrapped, @@ -234,6 +236,16 @@ def set_shared_experts(self, shared_experts): self.shared_experts = shared_experts self.use_nccl_stream = True + def get_expert_zero_copy_buffers(self): + """Buffers the experts should write their output / grad input into, if any. + + Returns: + A ``(output_buffer, grad_input_buffer)`` tuple. ``(None, None)`` unless the + dispatcher supports zero-copy, in which case the experts write straight into + the communication buffers instead of into fresh allocations. + """ + return None, None + class MoEAllGatherTokenDispatcher(MoETokenDispatcher): """ @@ -1598,6 +1610,20 @@ class _NCCLEPManager(_DispatchManager): lazily on the first dispatch, when the local token count is known. """ + # Zero-copy shared symm buffers, allocated once and reused across all layers/microbatches + # (class-level so every per-layer manager shares one set). _zc_fwd_token_buf is the forward symm + # buffer combine reads (the fc2 output); _zc_bwd_token_buf holds the backward grad. + # - fp8/fp4 (mxfp8 CuTe DSL grouped GEMM, Blackwell+): recv_tokens dies after FC1 quantizes it, + # so _zc_fwd_token_buf doubles as the dispatch recv_tokens; mcore also holds the dispatch + # probs (_zc_recv_topk_weights_buf) -- TE allocates nothing. + # - bf16 (op-fuser GroupedLinear GEMM, Hopper+): recv_tokens is the saved activation and + # can't double-duty, so TE pools the per-call recv_tokens/topk; _zc_fwd_token_buf holds only + # the fc2 output. + # TODO: move all to TE pool based allocation when symm memory pool supports cuda graph + _zc_fwd_token_buf = None + _zc_bwd_token_buf = None + _zc_recv_topk_weights_buf = None + def __init__( self, group: torch.distributed.ProcessGroup, @@ -1630,29 +1656,36 @@ def __init__( self.alignment = get_align_size_for_quantization(config) self.rank_capacity_factor = config.moe_expert_rank_capacity_factor self.static_shape = config.moe_ncclep_static_shape - if config.moe_ncclep_use_symm_mem: - raise NotImplementedError( - "moe_ncclep_use_symm_mem (symm-mem / zero-copy EP payload buffers) is not " - "supported yet." + self.zero_copy = config.moe_ncclep_zero_copy + self._zc_quant = self.zero_copy and bool(config.fp8 or config.fp4) + if self.zero_copy and not self.static_shape: + raise ValueError( + "moe_ncclep_zero_copy requires moe_ncclep_static_shape " + "(fixed [recv_capacity, hidden] symm buffers)." ) if self.static_shape: - if torch.cuda.get_device_capability()[0] < 10: + # static shape needs a fused grouped GEMM that consumes ragged per-expert counts on + # device (no host-side split narrowing): moe_grouped_gemm selects the grouped experts + # and use_transformer_engine_op_fuser fuses FC1+act+FC2 over them (fp8/fp4 via the CuTe + # DSL fused grouped MLP, bf16 via the op-fuser GroupedLinear grouped-tensor path). + if not (config.use_transformer_engine_op_fuser and config.moe_grouped_gemm): raise ValueError( - "moe_ncclep_static_shape=True requires an sm100+ (Blackwell or later) GPU with " - "a CuTe DSL / device-offset grouped GEMM; leave it False (dynamic shape) on " - "older GPUs." - ) - if not (config.use_transformer_engine_op_fuser or config.moe_grouped_gemm): - raise ValueError( - "moe_ncclep_static_shape=True requires the fused grouped GEMM; enable " - "use_transformer_engine_op_fuser (or moe_grouped_gemm)." - ) - if int(os.environ.get("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0")) <= 0: - raise ValueError( - "moe_ncclep_static_shape=True requires the CuTe DSL grouped GEMM; set " - "NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 (the expert grouped GEMM must consume ragged " - "per-expert counts on device)." + "moe_ncclep_static_shape=True requires BOTH use_transformer_engine_op_fuser " + "and moe_grouped_gemm (the fused grouped GEMM over device-side " + "per-expert counts)." ) + if config.fp8 or config.fp4: + if torch.cuda.get_device_capability()[0] < 10: + raise ValueError( + "moe_ncclep_static_shape=True with fp8/fp4 requires an sm100+ (Blackwell+) " + "GPU for the CuTe DSL grouped GEMM; leave it False on older GPUs." + ) + if int(os.environ.get("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "0")) <= 0: + raise ValueError( + "moe_ncclep_static_shape=True with fp8/fp4 requires the CuTe DSL grouped " + "GEMM; set NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 (the expert grouped GEMM must " + "consume ragged per-expert counts on device)." + ) if nccl_ep_dispatch is None: raise ImportError( @@ -1718,8 +1751,36 @@ def _ensure_bootstrap(self): if self.config.moe_flex_dispatcher_num_sms is not None else 0 ), - zero_copy=False, + zero_copy=self.zero_copy, ) + if self.zero_copy and _NCCLEPManager._zc_bwd_token_buf is None: + # Allocate once, shared across all managers. These are all persistent. + if self.config.overlap_moe_expert_parallel_comm: + # The 1F1B overlap schedule detaches the dispatch input, so autograd hands the + # dispatch-backward a non-symm clone of the grad_input buffer + warnings.warn( + "moe_ncclep_zero_copy + overlap_moe_expert_parallel_comm (1F1B EP overlap): " + "dispatch-backward gradient is not symm-mem-backed under the overlap schedule, " + "so it is staged into a symm buffer with one extra copy per dispatch-backward.", + stacklevel=2, + ) + assert ( + not torch.cuda.is_current_stream_capturing() + ), "zero-copy symm buffers must be allocated before CUDA-graph capture" + rc, h = self._recv_capacity, self.hidden_dim + _NCCLEPManager._zc_bwd_token_buf = alloc_ep_symm_buffer( + (rc, h), torch.bfloat16, self.group + ) + # The forward buffer combine reads (fc2 output). fp8 also feeds it to dispatch as + # recv_tokens (dead after FC1 quantize, so it double-duties); bf16 uses it only for fc2. + _NCCLEPManager._zc_fwd_token_buf = alloc_ep_symm_buffer( + (rc, h), torch.bfloat16, self.group + ) + if self._zc_quant: + # fp8 also owns the dispatch probs buffer (bf16 pools it per-call in TE). + _NCCLEPManager._zc_recv_topk_weights_buf = alloc_ep_symm_buffer( + (rc,), torch.float32, self.group + ) self._bootstrapped = True def dispatch( @@ -1748,10 +1809,18 @@ def dispatch( # tokens_per_expert: [num_local_experts] # dispatched_probs: [recv_capacity_per_rank] recv_tokens, tokens_per_expert, dispatched_probs = nccl_ep_dispatch( - self._buffer, hidden_states, topk_idx, topk_weights + self._buffer, + hidden_states, + topk_idx, + topk_weights, + recv_tokens=_NCCLEPManager._zc_fwd_token_buf if self._zc_quant else None, + recv_topk_weights=_NCCLEPManager._zc_recv_topk_weights_buf, ) self.tokens_per_expert = tokens_per_expert.to(torch.int64) - self.dispatched_probs = dispatched_probs + # fp8 zero-copy: dispatched_probs aliases the recv_topk_weights symm buffer, which the + # next layer's dispatch reuses; copy it out so it stays valid through this layer's backward. + # bf16 gets a fresh per-call pool buffer (not shared), so no copy is needed. + self.dispatched_probs = dispatched_probs.clone() if self._zc_quant else dispatched_probs return recv_tokens def get_permuted_hidden_states_by_experts(self, hidden_states: torch.Tensor) -> torch.Tensor: @@ -1790,7 +1859,10 @@ def combine( ) -> torch.Tensor: # hidden_states: [recv_capacity_per_rank, H] -> [num_local_tokens, H] hidden_states = nccl_ep_combine( - self._buffer, hidden_states, num_local_tokens=self.num_local_tokens + self._buffer, + hidden_states, + num_local_tokens=self.num_local_tokens, + grad_out=_NCCLEPManager._zc_bwd_token_buf, ) # Drop the buffer; backward keeps handle_mem alive via save_for_backward. self._buffer = None @@ -1871,6 +1943,30 @@ def __init__( "Please set --moe-flex-dispatcher-backend to deepep, deepepv2, hybridep, or ncclep" ) + def get_expert_zero_copy_buffers(self): + """NCCL-EP zero-copy: ``(output_buffer, grad_input_buffer)`` — the shared symm buffers the + experts write the fc2 output / fc1 dgrad into, so combine (fwd) and dispatch (bwd) read and + scatter them one-sided. ``(None, None)`` for every other backend/mode. + + Returned detached: the op-fuser calls requires_grad_() on its output and returns it, + so handing it the persistent buffer would permanently mark the shared classvar as requiring + grad and break the next layer's reuse. The detached view shares storage (zero-copy intact). + """ + + def _detached(name): + buf = getattr(self._comm_manager, name, None) + return buf.detach() if buf is not None else None + + # output_buffer (fc2 out / combine in) = _zc_fwd_token_buf; grad_input_buffer (fc1 dgrad / + # dispatch-bwd scatter) = _zc_bwd_token_buf. + # Under 1F1B overlap, feeding a symm grad_input_buffer is wasted: the overlap schedule's + # AccumulateGrad clones the fc1 dgrad into a plain buffer anyway. Return None so the + # op-fuser writes a plain dgrad; + dispatch_grad_input = ( + None if self.config.overlap_moe_expert_parallel_comm else _detached("_zc_bwd_token_buf") + ) + return _detached("_zc_fwd_token_buf"), dispatch_grad_input + def _initialize_metadata(self, routing_map: torch.Tensor, probs: torch.Tensor) -> torch.Tensor: """ Initialize the routing map and probs to a unified format covering the TPxEP group. diff --git a/megatron/core/transformer/moe/token_dispatcher_inference.py b/megatron/core/transformer/moe/token_dispatcher_inference.py index 081497f734c..b08d88f2641 100644 --- a/megatron/core/transformer/moe/token_dispatcher_inference.py +++ b/megatron/core/transformer/moe/token_dispatcher_inference.py @@ -40,6 +40,7 @@ gather_from_sequence_parallel_region, reduce_scatter_to_sequence_parallel_region, ) +from megatron.core.transformer.moe.inference_routing_mask_kernel import mask_routing_padding from megatron.core.transformer.moe.shared_experts import SharedExpertMLP from megatron.core.transformer.moe.token_dispatcher import MoEAllGatherTokenDispatcher from megatron.core.transformer.transformer_config import TransformerConfig @@ -313,6 +314,12 @@ class NVLSAllGatherVDispatcher(InferenceAllGatherDispatcherBase): _step_metadata: Optional[torch.Tensor] = None # [3] int32 _per_rank_worst_case_token_count: int = 2048 # round_up_tokens(max_tokens) // tp_size + # [1] int32 view onto context.gpu_view.real_token_count. Fixed GPU address; + # written each step by the context's transfer_bookkeeping_to_gpu(). Holds the + # real (unpadded) local token count so the dispatcher can mask routing for + # CUDA-graph padding tokens. Wired once by the context after gpu_view init. + _real_token_count_tensor: Optional[torch.Tensor] = None + # ── Class-level symmetric buffer handles (allocated once at model init) ─────── # Dtypes: hidden=bf16, routing=int64, probs=fp32, rsv=fp32. _symm_agv_hidden: Optional[dict] = None # {"tensor": ..., "handle": ...} @@ -326,6 +333,31 @@ def _get_rsv_tensor(cls) -> Optional[torch.Tensor]: unpermute output directly into it, avoiding a copy before RSV.""" return cls._symm_rsv["tensor"] if cls._symm_rsv is not None else None + @classmethod + def set_real_token_count_tensor(cls, tensor: torch.Tensor) -> None: + """Bind the context's GPU real-token-count tensor on the dispatcher class. + + Called once by DynamicInferenceContext after gpu_view is initialised. + The tensor is a fixed-address int32[1] view whose value is refreshed + each step by transfer_bookkeeping_to_gpu(). + """ + cls._real_token_count_tensor = tensor + + @classmethod + def modify_real_token_count_for_mtp(cls, mtp_token_count: int) -> None: + """Override the routing-mask token count for an MTP forward. + + Each step the context publishes batch_dimensions.token_count into the + bound tensor. MTP forwards are request-count shaped, so the controller + calls this before an MTP forward to point the mask at the MTP row count + instead. + """ + assert cls._real_token_count_tensor is not None, ( + "real-token-count tensor not wired; DynamicInferenceContext must " + "call set_real_token_count_tensor first" + ) + cls._real_token_count_tensor.fill_(mtp_token_count) + @classmethod def _rank_token_offset(cls) -> torch.Tensor: return cls._step_metadata[1:2] @@ -343,6 +375,7 @@ def _delete_buffers(cls): cls._symm_agv_probs = None cls._symm_rsv = None cls._symm_metadata = None + cls._real_token_count_tensor = None @classmethod def allocate_buffers( @@ -466,6 +499,10 @@ def __init__( runs_metadata_sync=runs_metadata_sync, ) self.topk = config.moe_router_topk + # Rank inside pg_collection.tp — the *standard* TP group that SP shards + # the routing map along. Base class self.tp_rank is the expt_tp rank, + # which is not what we want for the SP padding offset. + self.sp_rank = get_pg_rank(pg_collection.tp) # Set in dispatch_preprocess; consumed by token_dispatch and token_combine. self._local_tokens: int = 0 # When shared_expert_overlap is enabled, the shared expert forward is launched @@ -524,6 +561,18 @@ def token_dispatch(self, hidden_states, probs): if self._runs_metadata_sync: self.update_metadata(hidden_states.shape[0]) + # Mask out CUDA-graph padding rows of the local routing map so the AGV + # propagates -1 into agv_r for those slots; padding tokens then route + # to no expert. _real_token_count_tensor is wired by the context and + # holds the *global* unpadded token count, so we pass self.sp_rank to + # shift local rows into the global frame for the comparison. When unset + # (standalone dispatcher use without a context) all rows are real, so + # skip the mask. + if self.__class__._real_token_count_tensor is not None: + mask_routing_padding( + self.routing_map, self.__class__._real_token_count_tensor, self.sp_rank + ) + agv_h = self.__class__._symm_agv_hidden agv_r = self.__class__._symm_agv_routing agv_p = self.__class__._symm_agv_probs diff --git a/megatron/core/transformer/multi_latent_attention.py b/megatron/core/transformer/multi_latent_attention.py index 439a0f0649e..5cabd8697be 100644 --- a/megatron/core/transformer/multi_latent_attention.py +++ b/megatron/core/transformer/multi_latent_attention.py @@ -37,6 +37,7 @@ ) from megatron.core.transformer.attention import Attention, LinearProjBuilder from megatron.core.transformer.enums import AttnMaskType +from megatron.core.transformer.mla_qk_norm_config import QKNormConfigResolver from megatron.core.transformer.spec_utils import ModuleSpec, build_module from megatron.core.transformer.torch_norm import LayerNormBuilder from megatron.core.transformer.transformer_config import MLATransformerConfig @@ -519,10 +520,14 @@ def __init__( name=name, ) + # Resolve which classes to use for Q and KV linear up projections and norms, based on + # QK-norm selection. + layer_classes = self._resolve_qk_norm_config(submodules) + if self.config.q_lora_rank is None: # Not projecting query self.linear_q_proj = build_module( - submodules.linear_q_proj, + layer_classes["linear_q_proj"], self.config.hidden_size, self.config.num_attention_heads * self.q_head_dim, config=self.config, @@ -569,7 +574,7 @@ def __init__( ) self.linear_q_up_proj = build_module( - submodules.linear_q_up_proj, + layer_classes["linear_q_up_proj"], self.config.q_lora_rank, self.config.num_attention_heads * self.q_head_dim, config=self.config, @@ -616,7 +621,7 @@ def __init__( ) self.linear_kv_up_proj = build_module( - submodules.linear_kv_up_proj, + layer_classes["linear_kv_up_proj"], self.config.kv_lora_rank, self.config.num_attention_heads * (self.config.qk_head_dim + self.config.v_head_dim), config=self.config, @@ -631,18 +636,24 @@ def __init__( ) if self.config.q_lora_rank is not None: - self.q_layernorm = submodules.q_layernorm( + self.q_layernorm = layer_classes["q_layernorm"]( hidden_size=self.config.q_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, ) - self.kv_layernorm = submodules.kv_layernorm( + self.kv_layernorm = layer_classes["kv_layernorm"]( hidden_size=self.config.kv_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, ) + def _resolve_qk_norm_config( + self, submodules + ) -> dict[str, ModuleSpec | type | LayerNormBuilder]: + """Resolve which Q/KV norm and up-projection implementations to build.""" + return QKNormConfigResolver(self.config, submodules).resolve() + def _qkv_down_projection(self, hidden_states): """Unfused q/kv down projection path.""" if self.config.q_lora_rank is not None: @@ -1279,6 +1290,9 @@ def __init__( "FusedMLASelfAttention requires q_lora_rank to be set; " "fallback to MLASelfAttention for q_lora_rank=None." ) + # Resolve which linear class to use for Q and KV up projections, + # based on QK-norm selection. + layer_classes = self._resolve_qk_norm_config(submodules) qkv_down_proj_kwargs = {} if submodules.linear_qkv_down_proj in [TELinear]: @@ -1314,7 +1328,7 @@ def __init__( ) self.linear_q_up_proj = build_module( - submodules.linear_q_up_proj, + layer_classes["linear_q_up_proj"], self.config.q_lora_rank, self.config.num_attention_heads * self.q_head_dim, config=self.config, @@ -1329,7 +1343,7 @@ def __init__( ) self.linear_kv_up_proj = build_module( - submodules.linear_kv_up_proj, + layer_classes["linear_kv_up_proj"], self.config.kv_lora_rank, self.config.num_attention_heads * (self.config.qk_head_dim + self.config.v_head_dim), config=self.config, @@ -1343,12 +1357,12 @@ def __init__( name=(name + ".linear_kv_up_proj") if name is not None else None, ) - self.q_layernorm = submodules.q_layernorm( + self.q_layernorm = layer_classes["q_layernorm"]( hidden_size=self.config.q_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, ) - self.kv_layernorm = submodules.kv_layernorm( + self.kv_layernorm = layer_classes["kv_layernorm"]( hidden_size=self.config.kv_lora_rank, config=self.config, eps=self.config.layernorm_epsilon, diff --git a/megatron/core/transformer/transformer_config.py b/megatron/core/transformer/transformer_config.py index e46278e4058..0171e1b0582 100644 --- a/megatron/core/transformer/transformer_config.py +++ b/megatron/core/transformer/transformer_config.py @@ -151,6 +151,16 @@ class TransformerConfig(ModelParallelConfig): If attention backend is local we use the local pytorch implementation in mcore. Users can specify exact backend by changing this config. """ + flash_attention_version: Optional[Literal[2, 3, 4]] = None + """Pin the FlashAttention generation (2, 3, or 4) used by both the training + (TransformerEngine) and inference (mcore dynamic-batching) attention paths. When + None, each path selects a version automatically based on what is installed. Pinning + is required for batch-invariant mode: the training-side logprob recompute and the + inference engine must run the same kernel, since different FlashAttention + generations use different tile sizes and softmax accumulation orders and therefore + differ bitwise. On the training side this is enforced via TransformerEngine's + NVTE_FLASH_ATTN_V2/V3/V4 selection environment variables.""" + softmax_scale: Optional[float] = None """Softmax scale for attention scaling.""" @@ -781,6 +791,9 @@ class TransformerConfig(ModelParallelConfig): for each individual sample. - "global_aux_loss": Load balancing loss calculated at global batch level. - "sinkhorn": Balancing algorithm used in S-BASE. + - "quantile_balancing": Dual coordinate-descent quantile balancing (QB). Load balance is + handled entirely by an internal per-expert bias update; auxiliary losses must be disabled + (`moe_aux_loss_coeff` = 0) when QB is selected. - "none": No load balancing. A list of strings can be provided to combine multiple aux-loss load balancing types. The default is "aux_loss". @@ -853,6 +866,12 @@ class TransformerConfig(ModelParallelConfig): and decreased for the experts with more assigned tokens. The default value 1e-3 is same as that used in DeepSeekV3.""" + moe_router_quantile_balancing_ema: float = 0.0 + """EMA coefficient for the quantile-balancing per-expert bias (`qb_beta`), used only when + `moe_router_load_balancing_type` is "quantile_balancing". At each global batch the bias is + updated as `qb_beta = ema * qb_beta + (1 - ema) * local_quantile`. The default 0.0 means + no memory: the bias is replaced by the latest global-batch quantile estimate each step.""" + moe_router_force_load_balancing: bool = False """[Experimental] Force load balancing with random logits for MoE router, supports naive topk and group-limited topk. This is an experimental feature and only for benchmark.""" @@ -1032,11 +1051,11 @@ class TransformerConfig(ModelParallelConfig): later); the dispatcher asserts this. On older GPUs leave it False (dynamic shape). Defaults to False (narrow to the received tokens).""" - moe_ncclep_use_symm_mem: bool = False + moe_ncclep_zero_copy: bool = False """For the 'ncclep' flex dispatcher: use the NCCL symmetric-memory zero-copy IO path (ep_bootstrap zero_copy + symm-mem-backed receive/combine buffers) instead of the default HBM - staged-copy path. NOT SUPPORTED YET -- the dispatcher rejects this if set; the cross-stream - reuse ordering for the persistent symm-mem buffer is not implemented. Leave False.""" + staged-copy path, saving one copy on the wire. Requires moe_ncclep_static_shape and the fused op + (use_transformer_engine_op_fuser). Defaults to False.""" moe_mlp_glu_interleave_size: Optional[int] = None """When set, GLU activations in the MoE grouped MLP layer will use a @@ -1101,7 +1120,9 @@ class TransformerConfig(ModelParallelConfig): more details, see: https://pytorch.org/docs/stable/generated/torch.Tensor.backward.html.""" cuda_graph_warmup_steps: int = 3 - """Number of warmup steps for CUDA graphs""" + """Number of warmup steps for CUDA graphs. Note: GTP (``gtp_weight_remat_size > 1``) forces a + minimum of 2 per-graph warmup steps regardless of this value, because the first warmup builds + the weight-prefetch chain and the second exercises the prefetch path before capture.""" external_cuda_graph: bool = False """DEPRECATED and replaced by cuda_graph_impl. @@ -3098,6 +3119,27 @@ def _scope_to_str(s): "moe_input_jitter_eps is not supported with graphed moe recomputation." ) + if ( + self.gtp_weight_remat_size > 1 + and self.cuda_graph_impl == "local" + and (self.fp8 is not None or self.fp4 is not None) + and self.moe_shared_expert_intermediate_size is not None + and not self.moe_shared_expert_overlap + and ( + full_cudagraph + or CudaGraphModule.moe in self.cuda_graph_modules + or CudaGraphModule.moe_router in self.cuda_graph_modules + ) + ): + assert "shared_experts" not in self.recompute_modules, ( + "GTP + local CUDA graphs that capture shared_experts " + "(moe_router/moe scope) cannot recompute it under fp8/fp4: " + "te_checkpoint requires .backward(), but the local fwd-graph " + "warmup uses .grad(). Drop 'shared_experts' from " + "--recompute-modules (GTP-shard + offload instead), or use " + "--cuda-graph-impl full_iteration." + ) + if self.fine_grained_activation_offloading: offload_modules = set(self.offload_modules or []) if self.cuda_graph_impl == "local": @@ -3207,9 +3249,11 @@ def _scope_to_str(s): assert ( not self.moe_shared_expert_overlap ), 'disable moe_shared_expert_overlap when enabling overlap_moe_expert_parallel_comm' - assert ( - self.mtp_num_layers is None or self.mtp_num_layers == 1 - ), 'MTP layernum only supports 1 when enabling overlap_moe_expert_parallel_comm.' + assert self.mtp_num_layers in ( + None, + 0, + 1, + ), 'MTP supports at most one layer when enabling overlap_moe_expert_parallel_comm.' # NCCL EP (ncclep flex backend) mirrors hybridep's comm/compute overlap, but a few # configs are not yet safe under the 1F1B split and are gated here. @@ -3378,10 +3422,33 @@ def _scope_to_str(s): "for inference_optimized transformer implementation." ) + if self.flash_attention_version is not None: + assert self.flash_attention_version in (2, 3, 4), ( + "flash_attention_version must be one of 2, 3, or 4, got " + f"{self.flash_attention_version}" + ) + if self.batch_invariant_mode: assert ( self.attention_backend == AttnBackend.flash - ), "Batch invariant mode only supports FlashAttention" + ), "Batch invariant mode only supports FlashAttention (--attention-backend flash)" + # The training (TransformerEngine) and inference attention paths must run + # the same FlashAttention kernel, so the version cannot be left to each + # path's autodetection. FlashAttention-2 is excluded because it does not + # expose the fixed num_splits schedule the batch-invariant kernels require. + assert self.flash_attention_version in (3, 4), ( + "Batch invariant mode requires --flash-attention-version 3 or 4 so the " + "training and inference attention paths run the same batch-invariant " + f"FlashAttention kernel (got {self.flash_attention_version})." + ) + # Context parallelism routes through TE's FA2 fwd/bwd kernels directly, which + # cannot be pinned to another version; dropout is not batch-invariant. + assert ( + self.context_parallel_size == 1 + ), "Batch invariant mode does not support context parallelism" + assert ( + self.attention_dropout == 0.0 + ), "Batch invariant mode does not support attention dropout" if self.cuda_graph_impl != "none" and ( self.sequence_packing_scheduler is not None or self.dynamic_context_parallel diff --git a/megatron/core/transformer/transformer_layer.py b/megatron/core/transformer/transformer_layer.py index ea714552464..6155f7f7f48 100644 --- a/megatron/core/transformer/transformer_layer.py +++ b/megatron/core/transformer/transformer_layer.py @@ -21,7 +21,7 @@ from megatron.core.inference.utils import InferenceMode from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.transformer.cuda_graphs import is_graph_capturing, is_graph_warmup, make_weakref +from megatron.core.transformer.cuda_graphs import is_graph_capturing from megatron.core.transformer.enums import ( AttnMaskType, CudaGraphModule, @@ -2465,7 +2465,7 @@ def __init__(self, *args, **kwargs): self.is_moe_layer = True self.use_partial_cudagraphs = False self.moe_layer_recompute = False - self.token_dispatcher_attrs = {} + self._local_cudagraph_attr_names = None super().__init__(*args, **kwargs) @@ -2557,11 +2557,24 @@ def _resolve_token_dispatcher_attr(self, attr_name: str) -> tuple[Any, str]: obj = getattr(obj, parent_name) return obj, leaf_attr_name or attr_name - def _restore_token_dispatcher_attrs(self): - for attr_name, attr in self.token_dispatcher_attrs.items(): + def _restore_token_dispatcher_attrs(self, attr_outputs): + assert len(attr_outputs) == len(self._local_cudagraph_attr_names) + for attr_name, attr in zip(self._local_cudagraph_attr_names, attr_outputs): obj, name = self._resolve_token_dispatcher_attr(attr_name) setattr(obj, name, attr) + def _get_token_dispatcher_attrs(self): + attr_names = [] + token_dispatcher_attr_outputs = [] + for attr_name in self.mlp.token_dispatcher.cudagraph_attrs: + obj, name = self._resolve_token_dispatcher_attr(attr_name) + attr = getattr(obj, name) + if torch.is_tensor(attr): + attr_names.append(attr_name) + token_dispatcher_attr_outputs.append(attr) + + return tuple(attr_names), token_dispatcher_attr_outputs + def _forward_mlp_router( self, hidden_states, padding_mask=None, input_ids=None, packed_seq_params=None ): @@ -2588,7 +2601,7 @@ def _forward_mlp_router( if self.config.fp32_residual_connection: residual = residual.float() - router_outputs = apply_module(self.mlp)( + hidden_states, probs, shared_expert_output = apply_module(self.mlp)( pre_mlp_layernorm_output, intermediate_tensors=(), padding_mask=padding_mask, @@ -2596,17 +2609,25 @@ def _forward_mlp_router( packed_seq_params=packed_seq_params, ) - if is_graph_capturing() and not is_graph_warmup(): - for attr_name in self.mlp.token_dispatcher.cudagraph_attrs: - obj, name = self._resolve_token_dispatcher_attr(attr_name) - attr = getattr(obj, name) - if torch.is_tensor(attr): - attr.is_from_global_mempool = True - self.token_dispatcher_attrs[attr_name] = attr + if self.use_partial_cudagraphs: + attr_names, token_dispatcher_attr_outputs = self._get_token_dispatcher_attrs() + if self._local_cudagraph_attr_names is None: + self._local_cudagraph_attr_names = attr_names + else: + assert attr_names == self._local_cudagraph_attr_names + else: + # For eager mode, no need to pass the token_dispatcher attributes + token_dispatcher_attr_outputs = [] - return residual, *router_outputs + return ( + residual, + hidden_states, + probs, + shared_expert_output, + *token_dispatcher_attr_outputs, + ) - def _forward_mlp_expert_compute(self, hidden_states, probs): + def _forward_mlp_expert_compute(self, hidden_states, probs, token_dispatcher_attr_outputs): """ Executes the actual computation of the experts. @@ -2615,12 +2636,9 @@ def _forward_mlp_expert_compute(self, hidden_states, probs): step runs eagerly between the router and postprocess graph replays. """ - # During partial CUDA graph replay, use the probs returned from the graph in order - # to retain the router autograd edge. Rebinding it to the live router output ensures - # the backward DDP hook of router.weight is properly triggered. - if '_comm_manager.token_probs' in self.token_dispatcher_attrs: - self.token_dispatcher_attrs['_comm_manager.token_probs'] = probs - self._restore_token_dispatcher_attrs() + if self.use_partial_cudagraphs: + # Restore the token dispatcher attrs returned on the router graph's output surface. + self._restore_token_dispatcher_attrs(token_dispatcher_attr_outputs) self.mlp.fwd_execution_map = "expert_compute" return apply_module(self.mlp)(None, intermediate_tensors=(hidden_states, probs)) @@ -2637,16 +2655,7 @@ def _forward_mlp_postprocess(self, residual, output, shared_expert_output, mlp_b self.mlp.fwd_execution_map = "postprocess" output = apply_module(self.mlp)(None, intermediate_tensors=(output, shared_expert_output)) - out = self._forward_post_mlp((output, mlp_bias), residual) - - if is_graph_capturing() and not is_graph_warmup(): - for attr_name, attr in self.token_dispatcher_attrs.items(): - weak_ref = make_weakref(attr, inplace=False) - self.token_dispatcher_attrs[attr_name] = weak_ref - obj, name = self._resolve_token_dispatcher_attr(attr_name) - setattr(obj, name, weak_ref) - - return out + return self._forward_post_mlp((output, mlp_bias), residual) def _forward_mlp( self, @@ -2677,21 +2686,30 @@ def _forward_mlp_partial_cudagraphs( input_ids=None, packed_seq_params=None, ): - residual, hidden_states, probs, shared_expert_output = self._forward_mlp_router( + router_outputs = self._forward_mlp_router( hidden_states, padding_mask=padding_mask, input_ids=input_ids, packed_seq_params=packed_seq_params, ) + ( + residual, + hidden_states, + probs, + shared_expert_output, + *token_dispatcher_attr_outputs, + ) = router_outputs # After the router graph replays, the captured .copy_() operations that update - # self.token_dispatcher_attrs via `_maybe_dtoh_and_synchronize` are queued on the - # current stream but may not have completed. Record an event after the router + # the returned dispatcher tensors via `_maybe_dtoh_and_synchronize` are queued on + # the current stream but may not have completed. Record an event after the router # graph and wait on it, so we block only until the router's D2H copies complete. self._router_dtoh_event.record() self._router_dtoh_event.synchronize() - expert_output, mlp_bias = self._forward_mlp_expert_compute(hidden_states, probs) + expert_output, mlp_bias = self._forward_mlp_expert_compute( + hidden_states, probs, token_dispatcher_attr_outputs + ) return self._forward_mlp_postprocess( residual, expert_output, shared_expert_output, mlp_bias ) diff --git a/megatron/core/transformer/utils.py b/megatron/core/transformer/utils.py index 2249c79a2bd..aee4e961b9e 100644 --- a/megatron/core/transformer/utils.py +++ b/megatron/core/transformer/utils.py @@ -132,6 +132,28 @@ def make_sharded_tensors_for_checkpoint( tp_group = get_tensor_model_parallel_group_if_none(tp_group) dp_cp_group = parallel_state.get_data_parallel_group(with_context_parallel=True) + # GTP-sharded weights need the GTP axis layered onto the TP/DP offsets. The GTP helper + # is a no-op for non-GTP state_dicts, but importing it eagerly would be circular, so + # gate on HAVE_GTP and the presence of a GTP param before delegating. + from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + + if HAVE_GTP: + from megatron.core.tensor_parallel.gtp_api import ( + is_gtp_param, + make_sharded_tensors_for_checkpoint_with_gtp_remat, + ) + + if any(is_gtp_param(t) for t in state_dict.values()): + return make_sharded_tensors_for_checkpoint_with_gtp_remat( + state_dict, + prefix, + tensor_parallel_layers_axis_map, + sharded_offsets, + extra_state_suffix=extra_state_suffix, + tp_group=tp_group, + dp_cp_group=dp_cp_group, + ) + sharded_state_dict = {} for layer_name in state_dict.keys(): tensor = state_dict[layer_name] diff --git a/megatron/core/utils.py b/megatron/core/utils.py index ac97916d21b..7f721fe4ec8 100644 --- a/megatron/core/utils.py +++ b/megatron/core/utils.py @@ -935,7 +935,10 @@ def check_param_hashes_across_dp_replicas( for params, local_param_hashes, all_gather_group in zip( [non_expert_params, expert_params], [local_non_expert_param_hashes, local_expert_param_hashes], - [parallel_state.get_data_parallel_group(), parallel_state.get_expert_data_parallel_group()], + [ + parallel_state.get_data_parallel_group(with_gtp_remat=False), + parallel_state.get_expert_data_parallel_group(with_gtp_remat=False), + ], ): # Collect per-parameter hashes across all ranks in group. assert len(params) == len(local_param_hashes) @@ -1015,8 +1018,11 @@ def make_tp_sharded_tensor_for_checkpoint( new_offsets.append((tp_axis + prepend_axis_num, tp_rank, tp_size)) - if HAVE_DTENSOR and isinstance(tensor, DTensor): - # TP + FSDP2 sharding + is_torch_fsdp2_param = ( + hasattr(tensor, "is_torch_fsdp2_param") and HAVE_DTENSOR and isinstance(tensor, DTensor) + ) + if is_torch_fsdp2_param: + # When using FSDP2, every DP shard is a main replica. dp_replica_id = 0 tensor = tensor._local_tensor @@ -1029,10 +1035,54 @@ def make_tp_sharded_tensor_for_checkpoint( # FSDP2 shards axis 0 and TP shards some other axis new_offsets.append((prepend_axis_num, dp_rank, dp_size)) + # GTP: a GTP param additionally shards out_features (axis 0) by 1/gtp_remat. Layer that + # split onto TP offset — mirrors make_sharded_tensors_for_checkpoint_with_gtp_remat so direct + # callers (e.g. VocabParallelEmbedding, which can't use that wrapper because it needs + # allow_shape_mismatch) still save GTP weights with correct global offsets/shape. + from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + + if HAVE_GTP: + from megatron.core.fp8_utils import is_float8tensor + from megatron.core.tensor_parallel.gtp_api import dequantize_gtp_native_fp8, is_gtp_param + + if is_gtp_param(tensor): + gtp_rank = get_pg_rank(tensor.group) + gtp_remat_size = get_pg_size(tensor.group) + if tp_axis == 0: + # same axis as TP → one composite axis-0 offset + new_offsets[0] = ( + prepend_axis_num, + tp_rank * gtp_remat_size + gtp_rank, + tp_size * gtp_remat_size, + ) + else: + # GTP shards axis 0, TP shards a different axis → add a separate axis-0 offset + new_offsets.append((prepend_axis_num, gtp_rank, gtp_remat_size)) + # Elect the writer over the gtp_remat-EXCLUDED DP group (its true replicas). + dp_replica_id = parallel_state.get_data_parallel_rank( + with_context_parallel=True, with_gtp_remat=False + ) + # Saved global is the padded shape when GTP padded out_features for alignment. + if getattr(tensor, "pad_length", 0): + kwargs.setdefault("allow_shape_mismatch", True) + # Native-FP8 GTP shard: the param IS a QuantizedTensor (reports a fake BF16 dtype + # over FP8 bytes). Dequantize to real BF16 so the checkpoint stores portable + # high-precision values, not raw FP8 bytes mislabeled as BF16. Offsets above were + # already read from the FP8 param's GTP attrs; shape is preserved by dequantize. + # (dequantize_gtp_native_fp8 restores the base FP8 class for the dequantize call — + # TE's tex.dequantize does not recognize the dynamic GTP_ subclass.) + if is_float8tensor(tensor): + fp8_param = tensor + tensor = dequantize_gtp_native_fp8(tensor) + # Backlink to the live FP8 param: optimizer sharded_state_dict matches params + # to model entries by id(entry.data), which this dequantized copy would break + # (see _backfill_gtp_sharded_param_map in optimizer.py). + tensor._gtp_dequant_src = fp8_param + if replica_id is None: replica_id = (0, 0, dp_replica_id) - return ShardedTensor.from_rank_offsets( + sharded_tensor = ShardedTensor.from_rank_offsets( key, tensor, *prepend_offsets, @@ -1041,6 +1091,11 @@ def make_tp_sharded_tensor_for_checkpoint( prepend_axis_num=prepend_axis_num, **kwargs, ) + if is_torch_fsdp2_param: + # Marker used downstream for FSDP2-related logic, such as TP-DP + # sharding / loading for non-trivial parameters like SwiGLU. + sharded_tensor.is_torch_fsdp2_param = is_torch_fsdp2_param + return sharded_tensor def make_sharded_tensor_for_checkpoint(tensor, key, prepend_offsets=(), replica_id=None, **kwargs): @@ -1058,6 +1113,18 @@ def make_sharded_tensor_for_checkpoint(tensor, key, prepend_offsets=(), replica_ - dp_cp_group: Data parallel + context parallel group (default: None, falls back to parallel_state) """ + # Sanity guard. + from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + + if HAVE_GTP: + from megatron.core.tensor_parallel.gtp_api import is_gtp_param + + assert not is_gtp_param(tensor), ( + f"GTP weight-remat param '{key}' reached make_sharded_tensor_for_checkpoint (the " + "replicated path); route GTP-sharded weights through " + "make_tp_sharded_tensor_for_checkpoint or make_sharded_tensors_for_checkpoint instead." + ) + # Pop group parameters from kwargs tp_group = kwargs.pop('tp_group', None) dp_cp_group = kwargs.pop('dp_cp_group', None) @@ -1082,16 +1149,20 @@ def make_sharded_tensor_for_checkpoint(tensor, key, prepend_offsets=(), replica_ dp_size = get_pg_size(dp_cp_group) dp_replica_id = get_pg_rank(dp_cp_group) - if HAVE_DTENSOR and isinstance(tensor, DTensor): - # FSDP2 sharding + is_torch_fsdp2_param = ( + hasattr(tensor, "is_torch_fsdp2_param") and HAVE_DTENSOR and isinstance(tensor, DTensor) + ) + if is_torch_fsdp2_param: + # When using FSDP2, every DP shard is a main replica. dp_replica_id = 0 tensor = get_full_tensor_if_necessary(tensor) + # Add FSDP sharding rank offsets. new_offsets.append((prepend_axis_num, dp_rank, dp_size)) if replica_id is None: replica_id = (0, get_pg_rank(tp_group), dp_replica_id) - return ShardedTensor.from_rank_offsets( + sharded_tensor = ShardedTensor.from_rank_offsets( key, tensor, *prepend_offsets, @@ -1100,10 +1171,19 @@ def make_sharded_tensor_for_checkpoint(tensor, key, prepend_offsets=(), replica_ prepend_axis_num=prepend_axis_num, **kwargs, ) + if is_torch_fsdp2_param: + # Marker used downstream for FSDP2-related logic, such as TP-DP + # sharding / loading for non-trivial parameters like SwiGLU. + sharded_tensor.is_torch_fsdp2_param = is_torch_fsdp2_param + return sharded_tensor def get_full_tensor_if_necessary(tensor): - """For DTensor gets full tensor if some ranks will not have a local copy""" + """ + Captures an edge case where devices out-number elements in a DTensor, + for instance when generating a ShardedTensor. Replicate the DTensor + on all ranks to avoid empty DTensors on any rank. + """ need_full_tensor = False for i in range(tensor.device_mesh.ndim): if ( diff --git a/megatron/post_training/checkpointing.py b/megatron/post_training/checkpointing.py index 19df61b3b27..194b9bce100 100644 --- a/megatron/post_training/checkpointing.py +++ b/megatron/post_training/checkpointing.py @@ -17,7 +17,7 @@ from megatron.core import dist_checkpointing from megatron.core.dist_checkpointing.serialization import _legacy_common_state_exists -from megatron.core.utils import get_torch_version, is_torch_min_version, unwrap_model +from megatron.core.utils import unwrap_model from megatron.training import get_args from megatron.training.checkpointing import _load_base_checkpoint, load_checkpoint from megatron.training.utils import print_rank_0 @@ -231,3 +231,26 @@ def restore_sharded_modelopt_state(model: list[nn.Module], checkpoint_name: str model[0] = mto.restore_from_modelopt_state(model[0], common_modelopt_state) _load_extra_state_from_sharded_checkpoint(model[0], checkpoint_name, prefix="") + + +def load_kd_teacher_checkpoint(model) -> None: + """Load the teacher checkpoint for ModelOpt distillation if the model has one.""" + args = get_args() + if not getattr(args, "export_kd_teacher_load", None): + return + + teacher = unwrap_model(model[0]).teacher_model + print_rank_0( + f"Loading teacher as {type(teacher).__name__} from {args.export_kd_teacher_load} ..." + ) + # [WAR]: To avoid error out on loading teacher's checkpoint, we temporarily + # set args.finetune to True while loading the teacher checkpoint. + original_args_finetune, original_ckpt_format = args.finetune, args.ckpt_format + args.finetune = True + if args.export_kd_teacher_ckpt_format is not None: + args.ckpt_format = args.export_kd_teacher_ckpt_format + try: + load_checkpoint([teacher], None, None, load_arg='export_kd_teacher_load') + finally: + args.finetune, args.ckpt_format = original_args_finetune, original_ckpt_format + print_rank_0("... teacher loaded successfully.") diff --git a/megatron/post_training/model_builder.py b/megatron/post_training/model_builder.py index 3e3aabe989d..3b7ed3b96a3 100644 --- a/megatron/post_training/model_builder.py +++ b/megatron/post_training/model_builder.py @@ -5,11 +5,11 @@ import logging import os from argparse import Namespace -from typing import Any, Dict +from dataclasses import dataclass +from typing import Any, ClassVar, Dict import modelopt.torch.distill as mtd import modelopt.torch.distill.plugins.megatron as mtd_mcore -import modelopt.torch.opt as mto import yaml from megatron.core.models.gpt import GPTModel as MCoreGPTModel @@ -18,15 +18,83 @@ get_gpt_heterogeneous_layer_spec, ) from megatron.core.models.hybrid.hybrid_model import HybridModel as MCoreHybridModel +from megatron.core.pipeline_parallel.utils import is_pp_first_stage, is_pp_last_stage from megatron.core.post_training.modelopt.gpt.model_specs import get_gpt_modelopt_spec from megatron.core.post_training.modelopt.gpt.state_dict_hooks import ( mcore_gpt_load_te_state_dict_pre_hook, ) from megatron.core.post_training.modelopt.hybrid.model_specs import get_hybrid_stack_modelopt_spec +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.transformer.module import MegatronModule from megatron.post_training.checkpointing import load_modelopt_state -from megatron.post_training.utils import print_distributed_quant_summary from megatron.training import get_args, print_rank_0 from megatron.training.arguments import core_transformer_config_from_args +from megatron.training.models.gpt import GPTModelBuilder, GPTModelConfig +from megatron.training.models.hybrid import HybridModelBuilder, HybridModelConfig + + +@dataclass(kw_only=True) +class ModelOptModelConfig(GPTModelConfig): + """Config for the legacy ModelOpt model construction path. + + Identical to `GPTModelConfig` except for `builder` - construction still goes + through `gpt_config_from_args`, only the resolved builder class differs, since + ModelOpt-enabled runs need `ModelOptGPTModelBuilder` instead of `GPTModelBuilder`. + """ + + builder: ClassVar[str] = "megatron.post_training.model_builder.ModelOptGPTModelBuilder" + + +@dataclass(kw_only=True) +class ModelOptHybridModelConfig(HybridModelConfig): + """Config for the legacy ModelOpt model construction path, for hybrid models. + + Identical to `HybridModelConfig` except for `builder` - construction still goes + through `hybrid_config_from_args`. + """ + + builder: ClassVar[str] = "megatron.post_training.model_builder.ModelOptHybridModelBuilder" + + +class _ModelOptBuilderMixin: + """Shared `build_model()` override for the legacy ModelOpt model construction path. + + `modelopt_gpt_hybrid_builder` dispatches on `args.export_model_type` internally, so + the same implementation covers both GPT and hybrid models - only the parent + `ModelBuilder` (and its `build_distributed_models()`) differs per config type, so + each gets its own concrete class below rather than sharing one tied to `GPTModelBuilder`. + """ + + def build_model( + self, + pg_collection: ProcessGroupCollection, + pre_process: bool | None = None, + post_process: bool | None = None, + vp_stage: int | None = None, + ) -> MegatronModule: + args = get_args() + if pre_process is None: + pre_process = is_pp_first_stage(pg_collection.pp) + if post_process is None: + post_process = is_pp_last_stage(pg_collection.pp) + return modelopt_gpt_hybrid_builder( + args, + pre_process, + post_process, + vp_stage, + pg_collection=pg_collection, + ) + + +class ModelOptGPTModelBuilder(_ModelOptBuilderMixin, GPTModelBuilder): + """ModelBuilder adapter for the legacy ModelOpt model construction path.""" + + +class ModelOptHybridModelBuilder(_ModelOptBuilderMixin, HybridModelBuilder): + """ModelBuilder adapter for the legacy ModelOpt model construction path (hybrid).""" + + +logger = logging.getLogger(__name__) def count_parameters_in_layer(model, layer_name): @@ -49,69 +117,40 @@ def _load_teacher_model_config(checkpoint_path: str) -> Namespace: """Reads teacher config from a file. The config provided, either in the teacher checkpoint dir or via `--export-kd-teacher-model-config`, - should specify (in NeMo yaml config format) any model architecture settings which differ from the main student model's. - This function will translate NeMo field names to MCore as needed. + should specify any model architecture settings which differ from the main student model's. + The field names should match those returned by get_args() and not TransformerConfig. """ - required_teacher_fields = ( - "num_layers", - "hidden_size", - "ffn_hidden_size", - "num_attention_heads", - ) - args = get_args() + if args.export_kd_teacher_model_config is not None: config_path = args.export_kd_teacher_model_config + if not os.path.exists(config_path): + raise FileNotFoundError(f"Teacher model-config file ({config_path}) not found.") else: config_path = os.path.join(checkpoint_path, "model_config.yaml") - if not os.path.exists(config_path): - raise FileNotFoundError( - f"Teacher model-config file {config_path} not found.\n" - "Teacher checkpoint dir must contain a NeMo-format config named 'model_config.yaml'" - " or provide it via --export-kd-teacher-model-config." - ) - with open(config_path) as f: - config = yaml.safe_load(f) - - if missing_keys := [k for k in required_teacher_fields if k not in config]: - raise ValueError( - f"Teacher model config file ({config_path}) missing the following required fields: {missing_keys}" - ) - - if "encoder_seq_length" in config: - config["seq_length"] = config["encoder_seq_length"] - if "bias" in config: - config["disable_bias_linear"] = not config["bias"] - if config.get("activation") == "swiglu": - config["swiglu"] = True - if config.get("position_embedding_type", False) is None: - config["use_rotary_position_embeddings"] = config["no_position_embedding"] = True - if "share_embeddings_and_output_weights" in config: - config["untie_embeddings_and_output_weights"] = not config[ - "share_embeddings_and_output_weights" - ] - if "tokenizer" in config: - config["tokenizer_type"] = config["tokenizer"]["type"] - config["tokenizer_model"] = config["tokenizer"]["model"] - if "masked_softmax_fusion" in config: - config["no_masked_softmax_fusion"] = not config["masked_softmax_fusion"] - if config.get("normalization") == "layernorm1p": - config["apply_layernorm_1p"] = True - if "precision" in config: - config[config["precision"]] = True - if "mcore_gpt" in config: - config["use_mcore_models"] = config["mcore_gpt"] - - args_dict = vars(get_args()).copy() - del args_dict["kv_channels"] # not recalculated if present - # Setting teacher Flextron fields to false if training with Flextron, can be overridden - if "flextron" in args_dict: - config["flextron"] = False - if "enable_router" in args_dict: - config["enable_router"] = False - if "freeze_model" in args_dict: - config["freeze_model"] = False - args_dict.update(config) + if not os.path.exists(config_path): + logger.warning( + "No teacher config provided via --export-kd-teacher-model-config nor found at" + f" {checkpoint_path}/model_config.yaml. Assuming teacher model architecture same as student's." + ) # Useful for cases like QAD + config_path = None + + args_dict = vars(args).copy() + + if config_path is not None: + with open(config_path) as f: + config = yaml.safe_load(f) + + del args_dict["kv_channels"] # not recalculated if present + # Setting teacher Flextron fields to false if training with Flextron, can be overridden + if "flextron" in args_dict: + args_dict["flextron"] = False + if "enable_router" in args_dict: + args_dict["enable_router"] = False + if "freeze_model" in args_dict: + args_dict["freeze_model"] = False + + args_dict.update(config) # Backward compat: old checkpoints have hybrid_override_pattern but not hybrid_layer_pattern if ( @@ -156,7 +195,7 @@ def _build_teacher_model( _add_load_convert_hooks(teacher) - # NOTE: Checkpoint loading now handled in `megatron/training/checkpointing.py`. + # NOTE: Checkpoint loading now handled by `megatron.post_training.checkpointing.load_kd_teacher_checkpoint()`. return teacher @@ -374,7 +413,7 @@ def modelopt_gpt_hybrid_builder( ) if args.export_default_te_spec and args.export_te_mcore_model: - logging.getLogger(__name__).warning( + logger.warning( "--export-default-te-spec and --export-te-mcore-model are mutually exclusive. " "Since --export-default-te-spec is given, --export-te-mcore-model will be disabled." ) @@ -476,12 +515,7 @@ def modelopt_gpt_hybrid_builder( # Additional tweaks needed for MCore. # (accounts for sharded state, pipeline parallel, and potentially skipping LM loss) mtd_mcore.adjust_distillation_model_for_mcore(model, distill_cfg) - # Also remove KD mode state to prevent issues with re-conversion after restore. - mto.ModeloptStateManager( - model - ).state_dict().pop() # TODO(aanoosheh): remove once fixed in ModelOpt - print_distributed_quant_summary(model) return model diff --git a/megatron/post_training/utils.py b/megatron/post_training/utils.py index 5008e1b33dc..46cf2d696d4 100644 --- a/megatron/post_training/utils.py +++ b/megatron/post_training/utils.py @@ -9,8 +9,25 @@ from modelopt.torch.quantization.utils import is_quantized from packaging.version import Version -from megatron.core import parallel_state -from megatron.core.utils import unwrap_model + +def maybe_enable_modelopt(args): + """Set `args.modelopt_enabled` if a ModelOpt checkpoint or distillation teacher is + configured. Idempotent and safe to call multiple times (e.g. once early in + `pretrain_gpt.py` before building the model config, and again as a fallback in + `training.py` for callers that don't go through that entrypoint). + """ + if getattr(args, "modelopt_enabled", False): + return + + from megatron.post_training.checkpointing import has_modelopt_state + from megatron.training import print_rank_0 + + if args.load is not None and has_modelopt_state(args.load): + print_rank_0("ModelOpt checkpoint detected") + args.modelopt_enabled = True + if getattr(args, "export_kd_teacher_load", None): + # For distillation ckpts without ModelOpt state + args.modelopt_enabled = True def modelopt_version_higher_than(target_version: str): diff --git a/megatron/rl/agent/api.py b/megatron/rl/agent/api.py index e9fafaf162a..74b3783d2e7 100644 --- a/megatron/rl/agent/api.py +++ b/megatron/rl/agent/api.py @@ -19,12 +19,7 @@ LLMChatMessage, ReturnsRaw, ) -from ..rollout_granularity import ( - RELEASE_STATE_BY_SUBMISSION, - ConsumptionGranularity, - ReleaseState, - SubmissionGranularity, -) +from ..rollout_granularity import ConsumptionGranularity, SubmissionGranularity class AgentBaseModel(BaseModel, extra='allow'): @@ -102,14 +97,23 @@ def __getitem__(self, idx): GroupedRollouts = list[RolloutGroup] +class EpisodeResult(NamedTuple): + """All per-turn responses of one (possibly multi-turn) episode plus the final conversation.""" + + responses: list[InferenceResponse] + conversation: list[LLMChatMessage] + + class GroupRolloutParams(NamedTuple): """Returned by agent.prepare_group_rollout. One instance is created per group call and reused for all rollouts in that group. + Every rollout is an episode: run_episode generates it (one or more turns), while + build_rollout turns the completed episode into a Rollout. """ - inference_request: InferenceRequest - build_rollout: Callable[[InferenceResponse], Awaitable[Rollout]] + run_episode: Callable[[], Awaitable[EpisodeResult]] + build_rollout: Callable[[EpisodeResult], Awaitable[Rollout]] class ContrastiveRollout(AgentBaseModel): @@ -227,12 +231,18 @@ def _validate(request: GroupedRolloutRequest) -> None: class _SubmissionGate: - """Gate capacity is measured in units of the configured submission granularity.""" + """Gate capacity is measured in units of the configured submission granularity. + + Each granularity has a single release point: R slots free when inference + completes, so the gate bounds engine concurrency in rollouts. G and B + slots free when the trainer consumes the group/batch, so the gate + enforces the --rl-generation-lag run-ahead cap in groups/batches + respectively. + """ def __init__(self, *, capacity: int, submission: SubmissionGranularity) -> None: self._sem = asyncio.Semaphore(capacity) self._submission = submission - self._release_on = RELEASE_STATE_BY_SUBMISSION[submission] self.capacity = capacity # Observability counters, updated only on the configured submission # granularity (the only path that touches the semaphore). `held` @@ -251,8 +261,8 @@ async def acquire_for(self, granularity: SubmissionGranularity) -> None: self.held += 1 self.acquire_calls += 1 - def release_after(self, state: ReleaseState) -> None: - if self._release_on == state: + def release_for(self, granularity: SubmissionGranularity) -> None: + if self._submission == granularity: self._sem.release() self.held -= 1 self.release_calls += 1 @@ -279,7 +289,7 @@ class _InferredItem(NamedTuple): """One rollout post-inference, flowing from infer to assemble.""" item: _InferWorkItem - response: InferenceResponse + episode: EpisodeResult inferred_at: float = 0.0 @@ -391,16 +401,14 @@ async def _infer_worker(self) -> None: @trace_async_exceptions(verbose=True) async def _infer_one(self, item: _InferWorkItem) -> None: - response = await self.agent.get_rollout_response( - self.request, item.params.inference_request - ) + episode = await item.params.run_episode() inferred_at = time.monotonic() - self.gate.release_after("inferred") + self.gate.release_for("R") if item.infer_dequeued_at: self.engine_dwell.append(inferred_at - item.infer_dequeued_at) self.inferred_count += 1 await self.assemble_queue.put( - _InferredItem(item=item, response=response, inferred_at=inferred_at) + _InferredItem(item=item, episode=episode, inferred_at=inferred_at) ) async def stage_assemble(self) -> None: @@ -422,15 +430,17 @@ async def stage_assemble(self) -> None: completed = pending.pop(inferred.item.group_id) completed.sort(key=lambda item: item.item.rollout_idx) rollouts = await asyncio.gather( - *[item.item.params.build_rollout(item.response) for item in completed] + *[item.item.params.build_rollout(item.episode) for item in completed] ) - self.gate.release_after("assembled") self.assembled_count += 1 # NOTE: this filter is currently non-functional dead code: # _GranularityConfig._validate rejects filter_groups_with_same_reward # at pipeline construction, so `keep` is always True. Kept for a # future PR that regenerates dropped groups instead of - # under-delivering to the caller. + # under-delivering to the caller. That PR must also release the + # gate slot on the drop path: G/B slots free on consumption, and + # a dropped group never reaches stage_consume, so its slot (and + # eventually its batch's) would leak permanently. keep = ( not self.request.filter_groups_with_same_reward or np.std([rollout.reward for rollout in rollouts]) > 1e-6 @@ -468,6 +478,7 @@ async def stage_consume(self) -> AsyncIterator[RolloutGroup]: return self._record_output_dwell(group) yield group + self.gate.release_for("G") next_batch_id = 0 pending = self._consume_pending @@ -484,7 +495,8 @@ async def stage_consume(self) -> AsyncIterator[RolloutGroup]: next_batch_id += 1 for group in batch: yield group - self.gate.release_after("consumed") + self.gate.release_for("G") + self.gate.release_for("B") class GroupedRolloutGenerator(Agent, ABC): @@ -499,12 +511,7 @@ def __init__(self, *, parallel_generation_tasks: int | None = None, **kwargs): @abstractmethod async def prepare_group_rollout(self, request: GroupedRolloutRequest) -> GroupRolloutParams: - """Return the params for one group's rollouts. - - Called once per group by _RolloutPipeline.stage_prepare. The returned - build_rollout closure is invoked once per inference response in - _RolloutPipeline.stage_assemble. - """ + """Return the params for one group's rollouts.""" ... async def get_grouped_rollouts( diff --git a/megatron/rl/agent/reward_only_agent.py b/megatron/rl/agent/reward_only_agent.py index 972adc8a986..655a69cd22b 100644 --- a/megatron/rl/agent/reward_only_agent.py +++ b/megatron/rl/agent/reward_only_agent.py @@ -15,6 +15,7 @@ ReturnsTokens, ) from .api import ( + EpisodeResult, EvaluationAgent, EvaluationRequest, EvaluationResponse, @@ -41,6 +42,7 @@ class RewardOnlyAgent(RolloutGenerator, GroupedRolloutGenerator, PassAtEvaluatio """Agent that returns rollouts generated via default inference with a fixed reward function.""" env_id: str | None = None + max_turns: int = 1 def get_dataset(self, validation: bool = False): """Return validation or train dataset.""" @@ -84,49 +86,126 @@ def _get_rank_subset( return prompts[start_idx:end_idx] - async def _rollout_from_response( + async def get_observation( self, - request: RolloutRequest | GroupedRolloutRequest, + turn_idx: int, response: InferenceResponse, + conversation: list[LLMChatMessage], golden: Any, - ) -> Rollout: - assert isinstance( - request.inference_interface, ReturnsRaw - ), "InferenceInterface must support raw_text return to provide rollouts." - raw_text = response.raw_text + ) -> tuple[str | None, bool]: + """Return (observation, done) after a generation turn. Skipped on the last turn. - response_text = response.response.content + Override to implement multi-turn interactions. Must not mutate `conversation` or `golden`! + + Args: + turn_idx: 0-based index of the turn that just completed. + response: The inference response for this turn. + conversation: Message history before this turn's response was appended. + golden: Ground-truth / task data for reward computation. + + Returns: + (observation, done): If done is True the episode ends; observation is ignored. + If done is False, observation is a non-empty string that becomes the next user message. + """ + return None, True + + async def get_trajectory_reward( + self, responses: list[InferenceResponse], conversation: list[LLMChatMessage], golden: Any + ) -> float: + """Compute a scalar reward for the full trajectory. + + Override for trajectory-level or per-turn accumulated rewards. + """ + return await self.get_reward( + responses[-1].response.content, golden, responses[-1].finish_reason + ) + + async def _run_episode( + self, + request: RolloutRequest | GroupedRolloutRequest, + *, + prompt: str | list[LLMChatMessage], + golden: Any, + ) -> EpisodeResult: + """Run one (possibly multi-turn) episode over the group's prompt. + + Every turn takes the same path: prepare_request() on the conversation so far, then + get_rollout_response(). + get_observation() is consulted only while another generation is still possible; + on continue, the reply and observation are appended. + + Runs inside the infer stage, holding one submission slot for the whole episode. + """ + conversation = prompt + responses: list[InferenceResponse] = [] + + for turn_idx in range(self.max_turns): + turn_request = request.inference_interface.prepare_request( + conversation, request.generation_args + ) + # Adopt the request's prompt as the conversation: turn 0 may start from a bare + # string, which prepare_request normalizes into a single user message. + conversation = list(turn_request.prompt) + + response = await self.get_rollout_response(request, turn_request) + responses.append(response) + + if turn_idx + 1 < self.max_turns: + observation, done = await self.get_observation( + turn_idx, response, conversation, golden + ) + if done: + break + if not observation: + raise ValueError("get_observation must return a non-empty observation") + + conversation += [ + response.response, + LLMChatMessage(role="user", content=observation), + ] + + # The loop appends a reply only when continuing, so the final turn's reply is not in + # `conversation` yet; append it once so get_trajectory_reward sees the full dialogue. + return EpisodeResult( + responses=responses, conversation=conversation + [responses[-1].response] + ) + + async def _rollout_from_episode( + self, request: RolloutRequest | GroupedRolloutRequest, episode: EpisodeResult, golden: Any + ) -> Rollout | TokenRollout: + """Package a completed episode into a single rollout, one trajectory entry per turn. + + Calls `get_trajectory_reward()` once over all of the episode's responses. + """ + responses = episode.responses + reward = await self.get_trajectory_reward(responses, episode.conversation, golden) + problem_id = golden['problem_id'] if 'problem_id' in golden else None if isinstance(request.inference_interface, ReturnsTokens): - logprobs = response.logprobs - generation_mask = [ - True if (x >= response.prompt_length) else False - for x in range(len(response.token_ids)) - ] - rollout = TokenRollout( - trajectory=[response.token_ids], - reward=await self.get_reward(response_text, golden, response.finish_reason), - logprobs=[logprobs], - generation_mask=[generation_mask], + return TokenRollout( + trajectory=[r.token_ids for r in responses], + reward=reward, + logprobs=[r.logprobs for r in responses], + generation_mask=[ + [x >= r.prompt_length for x in range(len(r.token_ids))] for r in responses + ], env_id=self.env_id, - problem_id=golden['problem_id'] if 'problem_id' in golden else None, - policy_epoch=[response.policy_epoch], - kv_cache_epoch=[response.kv_cache_epoch], - num_evictions=[response.num_evictions], + problem_id=problem_id, + policy_epoch=[r.policy_epoch for r in responses], + kv_cache_epoch=[r.kv_cache_epoch for r in responses], + num_evictions=[r.num_evictions for r in responses], ) else: - rollout = Rollout( - trajectory=[raw_text], - reward=await self.get_reward(response_text, golden, response.finish_reason), + return Rollout( + trajectory=[r.raw_text for r in responses], + reward=reward, env_id=self.env_id, - problem_id=golden['problem_id'] if 'problem_id' in golden else None, - policy_epoch=[response.policy_epoch], - kv_cache_epoch=[response.kv_cache_epoch], - num_evictions=[response.num_evictions], + problem_id=problem_id, + policy_epoch=[r.policy_epoch for r in responses], + kv_cache_epoch=[r.kv_cache_epoch for r in responses], + num_evictions=[r.num_evictions for r in responses], ) - return rollout - async def get_rollout_response( self, request: RolloutRequest | GroupedRolloutRequest | EvaluationRequest, @@ -141,8 +220,7 @@ async def get_reward_rollouts(self, request: RolloutRequest) -> list[Rollout]: async def _single_rollout() -> Rollout: params = await self.prepare_group_rollout(request) - response = await self.get_rollout_response(request, params.inference_request) - return await params.build_rollout(response) + return await params.build_rollout(await params.run_episode()) return list(await asyncio.gather(*[_single_rollout() for _ in range(request.num_rollouts)])) @@ -150,13 +228,10 @@ async def prepare_group_rollout(self, request: GroupedRolloutRequest) -> GroupRo prompt, golden = await self.get_prompt(validation=request.validation) - inference_request = request.inference_interface.prepare_request( - prompt, request.generation_args - ) - + # Every rollout runs as a (possibly multi-turn) episode over the group's shared prompt. return GroupRolloutParams( - inference_request=inference_request, - build_rollout=functools.partial(self._rollout_from_response, request, golden=golden), + run_episode=functools.partial(self._run_episode, request, prompt=prompt, golden=golden), + build_rollout=functools.partial(self._rollout_from_episode, request, golden=golden), ) async def _evaluation( diff --git a/megatron/rl/inference/megatron.py b/megatron/rl/inference/megatron.py index e865a443c05..3348969da96 100644 --- a/megatron/rl/inference/megatron.py +++ b/megatron/rl/inference/megatron.py @@ -67,6 +67,12 @@ async def base_generate(self, request: InferenceRequest) -> InferenceResponse: extra_body={ "skip_prompt_log_probs": True, "add_BOS": (not args.rl_skip_bos_token and tokenizer.bos is not None), + # TODO: These are non-standard fields that add significant memory overheads to the + # chat completions payload. return_raw_text also wastes a lot of CPU cycles + # detokenizing prompt tokens, especially expensive for long prompts in agentic RL. + # Set to False if not needed in MRL. + "return_tokenized_data": True, + "return_raw_text": True, }, ) @@ -75,11 +81,11 @@ async def base_generate(self, request: InferenceRequest) -> InferenceResponse: return InferenceResponse( # TODO: Handle tool calls and reasoning in LLMChatMessage response=LLMChatMessage(**choice.message.model_dump(include={'role', 'content'})), - raw_text=choice.raw_text, - token_ids=choice.prompt_token_ids + choice.generation_token_ids, - logprobs=choice.generation_log_probs, + raw_text=choice.message.raw_text, + token_ids=choice.message.prompt_token_ids + choice.message.generation_token_ids, + logprobs=choice.message.generation_log_probs, finish_reason=choice.finish_reason, - prompt_length=len(choice.prompt_token_ids), + prompt_length=len(choice.message.prompt_token_ids), policy_epoch=choice.message.policy_epoch, kv_cache_epoch=choice.message.kv_cache_epoch, num_evictions=choice.message.num_evictions, diff --git a/megatron/rl/rl_utils.py b/megatron/rl/rl_utils.py index 3c3415e997d..8311d4640b6 100644 --- a/megatron/rl/rl_utils.py +++ b/megatron/rl/rl_utils.py @@ -65,7 +65,6 @@ RewardEvaluationResult, Rollout, RolloutGroup, - Rollouts, TokenRollout, ) from megatron.rl.agent.weighted_multi_task import WeightedMultiTask @@ -281,8 +280,8 @@ def verify_model_weights_swap( class RolloutStats: rewards: list[list[float]] # inner list is for a group env_ids: list[str] # same length as len(rewards) - turn_lens: list[list[int]] # token lengths of turns, grouped. - traj_lens: list[list[int]] # all turns comprise one trajectory. + turn_lens: list[list[int]] # tokens newly added by each turn, grouped. + traj_lens: list[list[int]] # the final sequence is the full trajectory. num_turns: None | list[list[int]] # num_turns per traj advantages: None | list[list[float]] min_piold_to_inf_prob: None | float @@ -511,7 +510,10 @@ def align_unpacked_inference_logprobs( # We need to align old_logprobs and inference logprobs as the latter are only for generations for i, inf_logprobs in enumerate(inference_logprobs): - first_gen_idx = first_gen_tok[i] + if not gen_masks_for_alignment[i].any(): + # No generation tokens; nothing to align. + continue + first_gen_idx = int(first_gen_tok[i]) # We subtract -1 here because we append eod token on the train side, and we do not # get it from the inference. For the eod token, we reuse old_logprobs value. end_idx = min(first_gen_idx + len(inf_logprobs), padded_inference_logprobs.shape[1]) @@ -917,26 +919,46 @@ def compute_group_stats( for turn_traj in rollout.trajectory: detokenized_traj = tokenizer.detokenize(turn_traj) lang_rl_log( - f"Rollout: [{rollout.env_id}] [{rollout.reward} : {len(rollout.trajectory)} tokens] {detokenized_traj}" + f"Rollout: [{rollout.env_id}] [{rollout.reward} : {len(turn_traj)} tokens] {detokenized_traj}" + ) + # A turn must never exceed the model's context window. + assert len(turn_traj) <= seq_len, ( + f"Rollout too long: {len(turn_traj)} > {seq_len} " + f"(last token {turn_traj[-1]})\n{detokenized_traj}" ) - # TODO(vitalyk): how does multiturn change EOD/EOT? - assert (len(turn_traj) == seq_len) or ( - turn_traj[-1] == tokenizer.eod - ), f"Rollout is not the correct length: {len(turn_traj)} {turn_traj[-1]}\n{detokenized_traj}" + # A single-turn completion can only end in eod or be truncated at seq_len. + # Multi-turn agents can additionally end a turn on a tool-call boundary. + if len(rollout.trajectory) == 1: + assert len(turn_traj) == seq_len or turn_traj[-1] == tokenizer.eod, ( + f"Single-turn rollout under seq_length must end in eod: " + f"len={len(turn_traj)} last={turn_traj[-1]}\n{detokenized_traj}" + ) else: lang_rl_log( - f"Rollout: [{rollout.env_id}] [{rollout.reward} : {len(rollout.trajectory)} chars] {rollout.trajectory}" + f"Rollout: [{rollout.env_id}] [{rollout.reward} : {len(rollout.trajectory)} turns] {rollout.trajectory}" ) group_num_turns.append(len(rollout.trajectory)) group_rewards.append(rollout.reward) - roll_turn_lens = [len(t) for t in rollout.trajectory] + # Multi-turn TokenRollout turns re-encode the full prior conversation. + # Report the incremental tokens each turn adds (env observation + generation); + # these sum to the final conversation length. + # Single-turn and raw (string) rollouts can use the plain per-turn length. + if isinstance(rollout, TokenRollout) and len(rollout.trajectory) > 1: + cumulative = [len(t) for t in rollout.trajectory] + roll_turn_lens = [cumulative[0]] + [ + cumulative[i] - cumulative[i - 1] for i in range(1, len(cumulative)) + ] + else: + roll_turn_lens = [len(t) for t in rollout.trajectory] group_turn_lengths.extend(roll_turn_lens) group_traj_lengths.append(sum(roll_turn_lens)) assert rollout.policy_epoch, "Rollout has no policy_epoch data" assert rollout.kv_cache_epoch, "Rollout has no kv_cache_epoch data" group_policy_epoch.append([epoch for turn in rollout.policy_epoch for _, epoch in turn]) group_kv_epoch.append([epoch for turn in rollout.kv_cache_epoch for _, epoch in turn]) - group_completed_epochs.extend(turn[-1][1] for turn in rollout.policy_epoch) + # completed_epochs is per-turn, so it cannot be masked per-rollout downstream + if rollout.trajectory: + group_completed_epochs.extend(turn[-1][1] for turn in rollout.policy_epoch) group_num_evictions.append(sum(rollout.num_evictions)) all_policy_epoch.append(group_policy_epoch) all_kv_cache_epoch.append(group_kv_epoch) @@ -994,12 +1016,17 @@ def prep_wandb_metrics( ): """Make a wandb-parseable dictionary of metrics for logging. + Zero-turn rollouts are a mark of placeholders (empty-trajectory pads for failed episodes). + Their 0.0 reward deliberately stays in the reward and group-mean/std aggregates: + it does affect training dynamics. + All other per-rollout field of theirs is masked out of stats. + Args: wandb_writer: Wandb run to log to. traj_lens: Grouped list of trajectory lengths. turn_lens: Grouped list of turn lengths. rewards: Grouped list of rewards. - num_turns: Grouped list of number of turns in the trajectories. + num_turns: Grouped list of number of turns in the trajectories. Zero means failure. advantages: Flattened list of advantages. policy_epoch: Grouped list of per-token policy epoch stamps. kv_cache_epoch: Grouped list of per-token KV cache epoch stamps. @@ -1009,22 +1036,47 @@ def prep_wandb_metrics( example_group: A list of rollouts of one group to log examples of trajectories. tokenizer: Tokenizer to untokenize trajectories for logging. """ + # Zero-turn rollouts are failure placeholders. + real_mask = [[nt > 0 for nt in g] for g in num_turns] + total_rollouts = sum(len(g) for g in num_turns) + failed_rollouts = sum(not keep for g in real_mask for keep in g) + failure_metrics = { + 'failed_rollouts/count': failed_rollouts, + 'failed_rollouts/ratio': (failed_rollouts / total_rollouts if total_rollouts else 0.0), + } + + def _real(grouped): + """Grouped per-rollout entries with placeholder (zero-turn) rollouts removed.""" + return [[x for x, keep in zip(g, m) if keep] for g, m in zip(grouped, real_mask)] + + # Reward metrics include failures. All other metrics do not. + table_rewards = [r for g in _real(rewards) for r in g] + traj_lens_real = _real(traj_lens) + num_turns_real = _real(num_turns) + policy_epoch_real = _real(policy_epoch) + kv_cache_epoch_real = _real(kv_cache_epoch) group_table = wandb_writer.Table( columns=['group_means', 'group_stds'], data=[[np.mean(g), np.std(g)] for g in rewards] ) # Per-rollout staleness (oldest token) - rollout_policy_staleness = [current_iteration - r[0] for g in policy_epoch for r in g] - rollout_kv_staleness = [current_iteration - r[0] for g in kv_cache_epoch for r in g] + rollout_policy_staleness = [current_iteration - r[0] for g in policy_epoch_real for r in g] + rollout_kv_staleness = [current_iteration - r[0] for g in kv_cache_epoch_real for r in g] # Per-rollout staleness (newest token) rollout_policy_last_token_staleness = [ - current_iteration - r[-1] for g in policy_epoch for r in g + current_iteration - r[-1] for g in policy_epoch_real for r in g + ] + rollout_kv_last_token_staleness = [ + current_iteration - r[-1] for g in kv_cache_epoch_real for r in g ] - rollout_kv_last_token_staleness = [current_iteration - r[-1] for g in kv_cache_epoch for r in g] # Per-token staleness - per_token_policy_staleness = [current_iteration - e for g in policy_epoch for r in g for e in r] - per_token_kv_staleness = [current_iteration - e for g in kv_cache_epoch for r in g for e in r] + per_token_policy_staleness = [ + current_iteration - e for g in policy_epoch_real for r in g for e in r + ] + per_token_kv_staleness = [ + current_iteration - e for g in kv_cache_epoch_real for r in g for e in r + ] metrics = { 'group_means_hist': wandb_writer.plot.histogram(group_table, 'group_means', 'Group Means'), @@ -1039,6 +1091,7 @@ def prep_wandb_metrics( 'advantages', 'Advantages', ), + # One row per real rollout. 'rollout_table': wandb_writer.Table( columns=[ 'reward', @@ -1051,9 +1104,9 @@ def prep_wandb_metrics( ], data=list( zip( - [r for g in rewards for r in g], - [l for g in traj_lens for l in g], - [e for g in num_evictions for e in g], + table_rewards, + [l for g in traj_lens_real for l in g], + [e for g in _real(num_evictions) for e in g], rollout_policy_staleness, rollout_kv_staleness, rollout_policy_last_token_staleness, @@ -1066,17 +1119,18 @@ def prep_wandb_metrics( columns=['policy_staleness', 'kv_staleness'], data=list(zip(per_token_policy_staleness, per_token_kv_staleness)), ), - 'mean_turn_length': np.mean([np.mean(g) for g in turn_lens]), - 'mean_turn_length_std': np.mean([np.std(g) for g in turn_lens]), - 'max_turn_length': max([max(g) for g in turn_lens]), - 'min_turn_length': min([min(g) for g in turn_lens]), - 'mean_traj_length': np.mean([np.mean(g) for g in traj_lens]), - 'mean_traj_length_std': np.mean([np.std(g) for g in traj_lens]), - 'max_traj_length': max([max(g) for g in traj_lens]), - 'min_traj_length': min([min(g) for g in traj_lens]), - 'mean_num_turns': np.mean([np.mean(g) for g in num_turns]), - 'max_num_turns': max([max(g) for g in num_turns]), - 'min_num_turns': min([min(g) for g in num_turns]), + # Group-level length/turn stats skip groups with all failed rollouts. + 'mean_turn_length': np.mean([np.mean(g) for g in turn_lens if g]), + 'mean_turn_length_std': np.mean([np.std(g) for g in turn_lens if g]), + 'max_turn_length': max(max(g) for g in turn_lens if g), + 'min_turn_length': min(min(g) for g in turn_lens if g), + 'mean_traj_length': np.mean([np.mean(g) for g in traj_lens_real if g]), + 'mean_traj_length_std': np.mean([np.std(g) for g in traj_lens_real if g]), + 'max_traj_length': max(max(g) for g in traj_lens_real if g), + 'min_traj_length': min(min(g) for g in traj_lens_real if g), + 'mean_num_turns': np.mean([np.mean(g) for g in num_turns_real if g]), + 'max_num_turns': max(max(g) for g in num_turns_real if g), + 'min_num_turns': min(min(g) for g in num_turns_real if g), 'mean_reward': np.mean([np.mean(g) for g in rewards]), 'mean_advantage': np.mean(advantages), 'nonzero_groups_ratio': np.count_nonzero(advantages) / len(advantages), @@ -1109,22 +1163,29 @@ def prep_wandb_metrics( 'staleness', 'Per-Token KV Cache Staleness', ), + **failure_metrics, } if example_group: if tokenizer is None: raise ValueError( "If you provide an example group to log, you need to provide a tokenizer too." ) + # Each turn in a trajectory is a cumulative sequence (prompt + all turns so far), so + # one row per rollout: the final turn already contains the whole conversation. metrics['rollouts'] = wandb_writer.Table( columns=['Trajectories', 'Tokens', 'Rewards'], rows=[ [ - tokenizer.detokenize(turn) if isinstance(r, TokenRollout) else turn, + ( + tokenizer.detokenize(r.trajectory[-1]) + if isinstance(r, TokenRollout) + else r.trajectory[-1] + ), r.trajectory, r.reward, ] for r in example_group - for turn in r.trajectory + if r.trajectory ], ) return metrics @@ -1328,21 +1389,27 @@ def maybe_log_training_metrics( wandb_writer.log(metrics, step=current_iteration) +PAD_TURN_UNIT = (None, -1) +"""Special filler entry to pad rollout turns for logprobs calculation. +Used to equalize per-rank trajectory counts without contributing to the loss.""" + + def prepare_trajectories( - rollouts: Rollouts, + rollout_turns: list[tuple[TokenRollout | Rollout | None, int]], tokenizer: MegatronTokenizer, seq_length: int, - sequence_packing: bool, skip_bos_token: bool, ): """Pad trajectories and extract the generation masks. + Args: - rollouts: Rollouts to extract trajectories from. + rollout_turns: (rollout, turn_idx) pairs; each pair becomes one trajectory. tokenizer: Tokenizer to get the padding token and potentially tokenize. seq_length: Maximum sequence length to pad to. Returns: - Trajectories and their generation masks. + Trajectories, their generation masks, and per-row inference logprobs + (unpadded tensor per real row, None per PAD row). Raises: ValueError: @@ -1383,43 +1450,47 @@ def prepare_trajectories( trajs = [] generation_masks = [] inference_logprobs = [] - for rollout in rollouts: - # traj, gen mask and logprobs are lists now. - # each list entry is a turn, single-turn environments just have a single-element list. - # We assume that all lengths of the structs above have the same lengths (number of turns). - - all_turns_trajectories = ( - copy.deepcopy(rollout.trajectory) + for rollout, turn_idx in rollout_turns: + if rollout is None: + # PAD_TURN_UNIT: inert filler so all DP ranks hold the same number of + # trajectories. All pad tokens, nothing generated, no inference logprobs. + trajs.append([tokenizer.pad] * seq_length) + generation_masks.append([False] * seq_length) + inference_logprobs.append(None) + continue + # traj, gen mask and logprobs are per-turn lists on the rollout; + # single-turn environments just have single-element lists. + # We assume that all the structs above have the same lengths (number of turns). + trajectory = ( + copy.deepcopy(rollout.trajectory[turn_idx]) if isinstance(rollout, TokenRollout) - else tokenizer.tokenize(rollout.trajectory) + else tokenizer.tokenize(rollout.trajectory)[turn_idx] ) - for turn_idx, trajectory in enumerate(all_turns_trajectories): - inf_logprobs = rollout.logprobs[turn_idx] - generation_mask = ( - rollout.generation_mask[turn_idx] if isinstance(rollout, TokenRollout) else None - ) - length = len(trajectory) - assert length <= seq_length, "Rollout too long, how did this happen?" - if len(trajectory) < seq_length: - assert ( - trajectory[-1] == tokenizer.eod - ), "Trajectories under a seq_length limit should have eod token at the end." - - if length < seq_length: - trajectory.extend([tokenizer.pad] * (seq_length - length)) - if generation_mask: - generation_mask.extend([False] * (seq_length - length)) - trajs.append(trajectory) - generation_masks.append(generation_mask) - - if inf_logprobs is not None: - inf_logprobs_tensor = torch.Tensor(inf_logprobs) - # Don't pad individual logprobs here - padding happens later if needed - inference_logprobs.append(inf_logprobs_tensor) - else: - inference_logprobs.append(None) + inf_logprobs = rollout.logprobs[turn_idx] + generation_mask = ( + copy.deepcopy(rollout.generation_mask[turn_idx]) + if isinstance(rollout, TokenRollout) + else None + ) + length = len(trajectory) + assert length <= seq_length, "Rollout too long, how did this happen?" + + if length < seq_length: + trajectory.extend([tokenizer.pad] * (seq_length - length)) + if generation_mask: + generation_mask.extend([False] * (seq_length - length)) + trajs.append(trajectory) + generation_masks.append(generation_mask) + + if inf_logprobs is not None: + inf_logprobs_tensor = torch.Tensor(inf_logprobs) + # Don't pad individual logprobs here - padding happens later if needed + inference_logprobs.append(inf_logprobs_tensor) + else: + inference_logprobs.append(None) - env_id_counts[rollout.env_id] += 1 + if turn_idx == 0: + env_id_counts[rollout.env_id] += 1 if torch.distributed.is_initialized(): logger.info(f"[{dist.get_rank()}] Rollout counts:") @@ -1429,20 +1500,14 @@ def prepare_trajectories( generation_masks = torch.tensor(generation_masks, dtype=torch.bool, device='cpu') trajs = torch.tensor(trajs, device='cpu') - # Only process if we have inference_logprobs - if inference_logprobs and any(lp is not None for lp in inference_logprobs): - # We need to pad all logprobs to the same size for sequence packing. - # For non-packing mode, keep as list of tensors (unpadded) - # This preserves the original behavior where each sequence can have different lengths - if sequence_packing: - inference_logprobs = _pad_nonnull_with_zeros(inference_logprobs, seq_length) - else: - inference_logprobs = None - - # Some sanity checks regarding the tokenization + # Some sanity checks regarding the tokenization. Pad units start with the pad + # token rather than bos, so the bos-equality check only applies to real rows. + real_rows = torch.tensor( + [rollout is not None for rollout, _ in rollout_turns], dtype=torch.bool + ) if not skip_bos_token: assert ( - tokenizer.bos is None or (trajs[:, 0] == tokenizer.bos).all() + tokenizer.bos is None or (trajs[real_rows][:, 0] == tokenizer.bos).all() ), "First token should be bos" else: assert ( @@ -1598,8 +1663,6 @@ def prepare_data_for_update( # We need this to correctly split the rollouts across dp groups. # And we do not actually need them grouped in anything below anyways. rollouts = [r for g in rollouts for r in g] - num_turns = [nt for g in group_stats.num_turns for nt in g] - total_turns_sampled = len(rollouts) # We might sample more than we consume in one step. samples_ratio_per_step = args.global_batch_size / ( @@ -1607,23 +1670,56 @@ def prepare_data_for_update( ) assert samples_ratio_per_step <= 1, "You cannot use more data than you sampled." - if (data_parallel_world_size := mpu.get_data_parallel_world_size()) > 0: - data_split_size = len(rollouts) // data_parallel_world_size + # Multi-turn rollouts contribute one trainable trajectory per turn, and turn counts vary. + # Flatten to single turns and split the turns across DP ranks. + + # advantages is already one entry per turn, so it is sliced with the same range. + rollout_turns = [ + (rollout, turn_idx) + for rollout in rollouts + for turn_idx in range(len(rollout.trajectory)) + ] + if not rollout_turns: + raise RuntimeError( + f"prepare_data_for_update: 0 usable trajectories from {len(rollouts)} rollout(s). " + "All rollouts have empty trajectories." + ) + + data_parallel_world_size = mpu.get_data_parallel_world_size() + # The total turn count is data-dependent, so it needs to be padded. + pad_to_multiple = data_parallel_world_size * args.micro_batch_size + if pad_n := -len(rollout_turns) % pad_to_multiple: + rollout_turns = rollout_turns + [PAD_TURN_UNIT] * pad_n + advantages = global_advantages = torch.cat( + [advantages, torch.zeros(pad_n, dtype=advantages.dtype, device=advantages.device)] + ) + total_turns_sampled = len(rollout_turns) + + has_inference_logprobs = any( + isinstance(rollout, TokenRollout) for rollout, _ in rollout_turns + ) + + if data_parallel_world_size > 0: + data_split_size = len(rollout_turns) // data_parallel_world_size data_split_range = ( mpu.get_data_parallel_rank() * data_split_size, (mpu.get_data_parallel_rank() + 1) * data_split_size, ) - rollouts = rollouts[data_split_range[0] : data_split_range[1]] - local_num_turns = sum(num_turns[data_split_range[0] : data_split_range[1]]) - steps_before = sum(num_turns[: data_split_range[0]]) - advantages = advantages[steps_before : steps_before + local_num_turns] + rollout_turns = rollout_turns[data_split_range[0] : data_split_range[1]] + advantages = advantages[data_split_range[0] : data_split_range[1]] # First we calculate them on a global level and then we split and recalculate on a local level. # Sequence packing and reporting needs it global but non-packing wants it local. with nvtx_range("rl/prepare-trajectories", time=True): trajs, generation_masks, inference_logprobs = prepare_trajectories( - rollouts, tokenizer, args.seq_length, sequence_packing, args.rl_skip_bos_token + rollout_turns, tokenizer, args.seq_length, args.rl_skip_bos_token ) + if not has_inference_logprobs: + inference_logprobs = None + elif sequence_packing: + # Pad each row to seq_length and stack; an all-PAD (all-None) local slice becomes an + # all-zero [num_rows, seq_length] tensor so this rank still joins the all_gather. + inference_logprobs = _pad_nonnull_with_zeros(inference_logprobs, args.seq_length) packing_context = None # Build trajectories based on sequence packing or standard processing @@ -1736,11 +1832,13 @@ def prepare_data_for_update( if inference_logprobs is not None: # Pack the inference logprobs using the helper function # We do this for logging purposes even if is_correction is disabled - packed_inference_logprobs = pack_inference_logprobs( - inference_logprobs=packing_context.original_inference_logprobs, - packing_info=packing_context.packing_info, - generation_masks=packing_context.original_generation_masks, - bin_size=args.seq_length, + packed_inference_logprobs, packed_inference_filled_mask = ( + pack_inference_logprobs( + inference_logprobs=packing_context.original_inference_logprobs, + packing_info=packing_context.packing_info, + generation_masks=packing_context.original_generation_masks, + bin_size=args.seq_length, + ) ) # Compute statistics for logging using packed data @@ -1749,14 +1847,11 @@ def prepare_data_for_update( packed_inference_logprobs=packed_inference_logprobs, packed_loss_mask=packing_context.packed_loss_mask, group_stats=group_stats, + filled_mask=packed_inference_filled_mask, ) # Store packed inference logprobs in packing context packing_context.packed_inference_logprobs = packed_inference_logprobs.cuda() - # Only mark as having inference logprobs for IS correction if enabled - packing_context.has_inference_logprobs = ( - args.rl_inference_logprobs_is_correction - ) with nvtx_range("rl/create-dataloader", time=True): # @vitalyk: This function also reconfigures the data loader to count the # global_batch_size in the bins frame of reference. @@ -2286,8 +2381,9 @@ def _pad_nonnull_with_zeros(data: list[Optional[torch.Tensor]], max_len: int) -> A padded tensor which is a stacked list of padded input tensors. """ - if all([el is None for el in data]): - raise ValueError("At least one element of the data list should be not None.") + if all(el is None for el in data): + # All rows are PAD; return an all-zero tensor so that no DP rank stalls. + return torch.zeros((len(data), max_len)) padded_data = [] for chunk in data: if chunk is not None: diff --git a/megatron/rl/rollout_granularity.py b/megatron/rl/rollout_granularity.py index 7bc13ab5b21..b9432f0bd4d 100644 --- a/megatron/rl/rollout_granularity.py +++ b/megatron/rl/rollout_granularity.py @@ -6,14 +6,6 @@ SubmissionGranularity = Literal["R", "G", "B"] ConsumptionGranularity = Literal["G", "B"] -ReleaseState = Literal["inferred", "assembled", "consumed"] - - -RELEASE_STATE_BY_SUBMISSION: dict[SubmissionGranularity, ReleaseState] = { - "R": "inferred", - "G": "assembled", - "B": "consumed", -} def get_rl_parallel_generation_tasks(args) -> int: diff --git a/megatron/rl/sequence_packing_utils.py b/megatron/rl/sequence_packing_utils.py index c78e7d39127..fe609b6b7f0 100644 --- a/megatron/rl/sequence_packing_utils.py +++ b/megatron/rl/sequence_packing_utils.py @@ -497,7 +497,10 @@ def pack_inference_logprobs( bin_size: Size of each bin Returns: - Packed inference logprobs tensor of shape [num_bins, bin_size - 1] + Packed inference logprobs tensor of shape [num_bins, bin_size - 1], and a + bool mask of the same shape marking positions actually filled from the + engine (the train side appends tokens, e.g. EOD, that the engine never + reported a logprob for; those stay zero-filled and must not be compared). """ num_bins = len(packing_info.bin_seq_indices) @@ -505,6 +508,7 @@ def pack_inference_logprobs( packed_inference_logprobs = torch.zeros( (num_bins, bin_size - 1), dtype=torch.float32, device='cpu' ) + filled_mask = torch.zeros((num_bins, bin_size - 1), dtype=torch.bool, device='cpu') # Create mapping from global sequence index to local bin index # This is needed because seq_to_bin_idx uses global bin indices, @@ -549,8 +553,9 @@ def pack_inference_logprobs( packed_inference_logprobs[local_bin_idx, pack_start:pack_end] = seq_inf_logprobs[ :actual_len ] + filled_mask[local_bin_idx, pack_start:pack_end] = True - return packed_inference_logprobs + return packed_inference_logprobs, filled_mask def compute_packed_inference_logprobs_stats( @@ -558,6 +563,7 @@ def compute_packed_inference_logprobs_stats( packed_inference_logprobs: torch.Tensor, packed_loss_mask: torch.Tensor, group_stats: Any, + filled_mask: Optional[torch.Tensor] = None, ) -> None: """Compute statistics for packed inference logprobs for logging purposes. @@ -569,17 +575,30 @@ def compute_packed_inference_logprobs_stats( packed_inference_logprobs: Packed inference logprobs [num_bins, seq_len-1] packed_loss_mask: Loss mask indicating valid positions [num_bins, seq_len] group_stats: Statistics object to update with computed metrics + filled_mask: Optional bool mask [num_bins, seq_len-1] marking positions actually + filled from the engine's reported logprobs. Positions the engine never + reported (e.g. the train-side EOD append) stay zero-filled and are excluded + from the stats so they do not contribute spurious |p_old - 1| terms. """ # Lazy import to avoid circular dependency (rl_utils imports from this module) from megatron.rl.rl_utils import update_inference_logprobs_group_stats - # Ensure all tensors are on the same device (CPU for stats computation) + # Ensure all tensors are on the same device (CPU for stats computation). + # Compare in the training dtype: old_logprobs are bf16 while the engine reports + # fp32 logprobs, so comparing raw values shows bf16-rounding noise even when the + # two sides are bitwise identical (the unpacked path already rounds by writing + # into an old_logprobs-dtype buffer in align_unpacked_inference_logprobs). old_logprobs = old_logprobs.cpu() - packed_inference_logprobs = packed_inference_logprobs.cpu() + packed_inference_logprobs = packed_inference_logprobs.cpu().to(old_logprobs.dtype) packed_loss_mask = packed_loss_mask.cpu() # Use packed_loss_mask to identify valid positions for stats (shift by 1 for logprobs) mask = packed_loss_mask[:, 1:].bool() + if filled_mask is not None: + # Exclude positions the engine never reported (zero-filled in packing), e.g. + # the train-side EOD append: comparing exp(0)=1 against a real prob there + # poisons the mismatch stats with spurious |p_old - 1| terms. + mask = mask & filled_mask.to(mask.device) # Ensure shapes match if mask.shape != old_logprobs.shape: diff --git a/megatron/training/argument_utils.py b/megatron/training/argument_utils.py index f6cee570828..a07ef3601a0 100644 --- a/megatron/training/argument_utils.py +++ b/megatron/training/argument_utils.py @@ -557,8 +557,18 @@ def _default_config_from_args(cls: type, args: Namespace, return_instance: bool return kwargs -def gpt_config_from_args(args: Namespace, config: TransformerConfig | None = None) -> Any: - """Create a GPTModelConfig from the appropriate values in the `args` Namespace.""" +def gpt_config_from_args( + args: Namespace, + config: TransformerConfig | None = None, + model_config_cls: type = GPTModelConfig, +) -> Any: + """Create a GPTModelConfig (or a compatible subclass) from the `args` Namespace. + + `model_config_cls` lets callers reuse this same arg-derivation logic for + subclasses that only override metadata (e.g. `builder`) and add no new fields, + such as `ModelOptModelConfig`. + """ + assert issubclass(model_config_cls, GPTModelConfig) kwargs = {} if config is None: @@ -599,11 +609,21 @@ def gpt_config_from_args(args: Namespace, config: TransformerConfig | None = Non kwargs["vocab_size"] = args.vocab_size kwargs["should_pad_vocab"] = True - return GPTModelConfig(**kwargs) + return model_config_cls(**kwargs) + +def hybrid_config_from_args( + args: Namespace, + config: TransformerConfig | None = None, + model_config_cls: type = HybridModelConfig, +) -> Any: + """Create a HybridModelConfig (or a compatible subclass) from the `args` Namespace. -def hybrid_config_from_args(args: Namespace, config: TransformerConfig | None = None) -> Any: - """Create a HybridModelConfig from the appropriate values in the `args` Namespace.""" + `model_config_cls` lets callers reuse this same arg-derivation logic for + subclasses that only override metadata (e.g. `builder`) and add no new fields, + such as `ModelOptHybridModelConfig`. + """ + assert issubclass(model_config_cls, HybridModelConfig) kwargs = {} if config is None: @@ -649,7 +669,7 @@ def hybrid_config_from_args(args: Namespace, config: TransformerConfig | None = kwargs["vocab_size"] = args.vocab_size kwargs["should_pad_vocab"] = True - return HybridModelConfig(**kwargs) + return model_config_cls(**kwargs) def pretrain_cfg_container_from_args(args: Namespace, model_cfg=None) -> PretrainConfigContainer: diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index a168381eaa7..16233ee9add 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -455,10 +455,23 @@ def validate_args(args, defaults={}): update_use_dist_ckpt(args) + # GTP_remat counts toward total_model_size (an independent weight-shard axis), so the + # args.data_parallel_size below is the replicate degree (matches + # parallel_state). gtp_weight_remat_size is derived from --tensor-parallel-num-weight-shards. + from megatron.core.model_parallel_config import resolve_tensor_parallel_weight_shards + + (args.tensor_parallel_num_weight_shards, args.gtp_weight_remat_size) = ( + resolve_tensor_parallel_weight_shards( + args.tensor_model_parallel_size, + args.tensor_parallel_num_weight_shards, + getattr(args, "gtp_weight_remat_size", 1), + ) + ) total_model_size = ( args.tensor_model_parallel_size * args.pipeline_model_parallel_size * args.context_parallel_size + * args.gtp_weight_remat_size ) # Total model size. @@ -478,6 +491,7 @@ def validate_args(args, defaults={}): args.tensor_model_parallel_size * args.pipeline_model_parallel_size * args.context_parallel_size + * args.gtp_weight_remat_size ) args.data_parallel_size = args.world_size // total_model_size @@ -946,7 +960,10 @@ def validate_args(args, defaults={}): ) # Infer use of MLA from unified pattern - if args.hybrid_layer_pattern and Symbols.DS_ATTENTION in args.hybrid_layer_pattern: + if args.hybrid_layer_pattern and ( + Symbols.MLA in args.hybrid_layer_pattern + or Symbols.DS_ATTENTION in args.hybrid_layer_pattern + ): args.multi_latent_attention = True # === End of hybrid layer pattern: deprecation handling and validation === @@ -1187,15 +1204,22 @@ def validate_args(args, defaults={}): ): raise ValueError("MXFP8 with inference optimized layers requires FlashInfer >= 0.6.4") - if args.inference_dynamic_batching_sampling_backend == 'flashinfer': - try: - import flashinfer # noqa: F401 - except ImportError as e: - raise ImportError( - "--inference-dynamic-batching-sampling-backend=flashinfer requires " - "the flashinfer package; install it or pass " - "--inference-dynamic-batching-sampling-backend=torch." - ) from e + # Streaming dequantize is unsafe with tensorwise (current) FP8 scaling. + # The streaming planner does a BF16->FP8 ``copy_`` per slice; tensorwise + # recomputes the per-tensor scale from each slice's amax, so multi-shard + # destinations (e.g. resharded loads) end up with inconsistent scales + # across slices and the loaded weights are corrupted. Block-scaled + # recipes (mxfp8, blockwise, nvfp4) carry per-block scales and are + # unaffected. Force the upfront ``force_all_tensors_to_non_fp8`` path + # for tensorwise. + if args.fp8 and args.fp8_recipe == "tensorwise" and args.stream_ckpt_dequant: + warn_rank_0( + "--fp8-recipe=tensorwise is incompatible with the streaming " + "checkpoint dequantize path; falling back to the upfront " + "dequantize pass. Pass --no-stream-ckpt-dequant to silence " + "this warning." + ) + args.stream_ckpt_dequant = False if args.use_megatron_fsdp: # NOTE: The flag `use_custom_fsdp` is deprecated and will be removed in future versions. @@ -1654,6 +1678,106 @@ def validate_args(args, defaults={}): if args.expert_model_parallel_size > 1 and 'ep_dp' not in args.high_priority_stream_groups: args.high_priority_stream_groups.append('ep_dp') + # Derive the internal gtp_weight_remat_size from the user-facing + # --tensor-parallel-num-weight-shards. gtp_weight_remat_size has no CLI flag (it is excluded + # from argument generation), so it is set here as a fresh attribute on args before it is + # consumed below (and in initialize/training, which read args.gtp_weight_remat_size directly). + # Mirrors ModelParallelConfig.__post_init__. + from megatron.core.model_parallel_config import resolve_tensor_parallel_weight_shards + + (args.tensor_parallel_num_weight_shards, args.gtp_weight_remat_size) = ( + resolve_tensor_parallel_weight_shards( + args.tensor_model_parallel_size, + args.tensor_parallel_num_weight_shards, + getattr(args, "gtp_weight_remat_size", 1), + ) + ) + # Same for the expert layers: derive the internal expert_gtp_weight_remat_size from the + # user-facing --expert-tensor-parallel-num-weight-shards (expert_tensor_parallel_size is + # defaulted earlier in validate_args). expert_gtp_weight_remat_size has no CLI flag. + (args.expert_tensor_parallel_num_weight_shards, args.expert_gtp_weight_remat_size) = ( + resolve_tensor_parallel_weight_shards( + args.expert_tensor_parallel_size, + args.expert_tensor_parallel_num_weight_shards, + getattr(args, "expert_gtp_weight_remat_size", 1), + ) + ) + + if args.gtp_weight_remat_size > 1 or args.expert_gtp_weight_remat_size > 1: + if args.fp4 and not args.fp4_param_gather: + raise ValueError( + "GTP (--tensor-parallel-num-weight-shards / " + "--expert-tensor-parallel-num-weight-shards > 1) with --fp4-format requires " + "--fp4-param-gather so NVFP4 weights are all-gathered as native NVFP4." + ) + gtp_weight_remat_size = args.gtp_weight_remat_size + egtp_weight_remat_size = args.expert_gtp_weight_remat_size + if get_device_arch_version() >= 10: + # Setting GTP communication groups for high priority streams for Blackwell and later + # architectures. Assigning high priority to communication streams ensures that + # communication kernels are scheduled with higher priority, minimizing the exposed + # communication when it is overlapped with other computation kernels. + if 'gtp_remat' not in args.high_priority_stream_groups: + args.high_priority_stream_groups.append('gtp_remat') + warn_rank_0("Setting 'gtp_remat' group for high priority streams.") + if ( + egtp_weight_remat_size > 1 + and 'expt_gtp_remat' not in args.high_priority_stream_groups + ): + args.high_priority_stream_groups.append('expt_gtp_remat') + warn_rank_0("Setting 'expt_gtp_remat' group for high priority streams.") + + # Sanity check for 'CUDA_GRAPHS_USE_NODE_PRIORITY'. + if args.cuda_graph_impl != "none": + assert os.environ.get('CUDA_GRAPHS_USE_NODE_PRIORITY') == "1", ( + 'GTP requires CUDA_GRAPHS_USE_NODE_PRIORITY=1 to make sure fine-grained GTP ' + 'comms can be well overlapped with GEMMs when CudaGraph is enabled for ' + 'Blackwell and later architecture.' + ) + + # Sanity check for 'NCCL_PROTO'. + if os.environ.get('NCCL_PROTO', '').lower() == "simple": + warn_rank_0( + "Generally GTP prefers 'NCCL_PROTO=LL128 or LL' while get 'NCCL_PROTO=simple', " + "force setting NCCL_PROTO=Simple might introduce bad perf." + ) + + assert not args.ddp_average_in_collective, ( + "GTP requires --ddp-average-in-collective off (the default); averaged collectives " + "would need per-buffer 1/gtp_remat scaling." + ) + + assert args.ckpt_format in ('torch', 'torch_dist'), ( + f"GTP supports only --ckpt-format 'torch' (legacy) or 'torch_dist', got " + f"'{args.ckpt_format}'." + ) + assert not ( + getattr(args, 'dist_ckpt_optim_fully_reshardable', False) + and getattr(args, 'distrib_optim_fully_reshardable_mem_efficient', False) + ), ( + "GTP does not support the distributed-optimizer fully-reshardable + " + "mem-efficient checkpoint mode. Disable " + "--distrib-optim-fully-reshardable-mem-efficient (or " + "--dist-ckpt-optim-fully-reshardable)." + ) + + # GTP with the mxfp8 recipe requires --fp8-param-gather: GTP keeps no bf16 weight and + # relies on the optimizer maintaining the fp8 shard (the forward all-gathers fp8 and does + # not re-quantize). Without fp8-param-gather the fp8 forward weight would never be updated. + if getattr(args, 'fp8_recipe', None) == 'mxfp8': + assert getattr(args, 'fp8_param_gather', False), ( + "GTP + mxfp8 requires --fp8-param-gather (the optimizer maintains the fp8 shard; " + "GTP does not keep or re-quantize a bf16 weight)." + ) + # MXFP8 params cannot be mapped into the contiguous param buffer (TE's + # replace_raw_data does not support the MXFP8 tile-scaling layout), so the param + # all-gather must reuse the grad buffer instead. + assert getattr(args, 'reuse_grad_buf_for_mxfp8_param_ag', False), ( + "GTP + mxfp8 + --fp8-param-gather requires --reuse-grad-buf-for-mxfp8-param-ag " + "(MXFP8 params keep their own quantized storage; mapping them into the param " + "buffer via replace_raw_data is unsupported)." + ) + # Disable bias gelu fusion if we are disabling bias altogether if not args.add_bias_linear: args.bias_gelu_fusion = False @@ -1978,6 +2102,39 @@ def validate_args(args, defaults={}): if not args.async_save: args.async_strategy = "mcore" + if args.logits_save_dir is not None: + assert ( + args.logits_save_top_k is not None + ), '--logits-save-top-k is required when --logits-save-dir is set.' + assert args.async_save, ( + '--logits-save-dir requires --async-save (and --use-persistent-ckpt-worker). ' + 'Logits are flushed as an async request in the checkpoint queue.' + ) + if not args.freeze_all_layers: + warn_rank_0( + '--logits-save-dir without --freeze-all-layers: the LM loss is still computed and ' + 'gradients will update the model while logits are dumped. This is intended only ' + 'when dumping logits during active training; for a frozen-teacher dump pass ' + '--freeze-all-layers.' + ) + + if args.freeze_all_layers: + if args.use_distributed_optimizer: + warn_rank_0( + '--freeze-all-layers incompatible with use_distributed_optimizer. Disabling use_distributed_optimizer.' + ) + args.use_distributed_optimizer = False + if args.overlap_param_gather: + warn_rank_0( + '--freeze-all-layers incompatible with overlap_param_gather. Disabling overlap_param_gather.' + ) + args.overlap_param_gather = False + + if args.override_ckpt_iteration is not None: + assert ( + not args.finetune + ), "Cannot override checkpoint iteration together with finetune flag." + # Inference args if args.inference_batch_times_seqlen_threshold > -1: assert ( @@ -2576,6 +2733,14 @@ def _add_inference_args(parser): 'Falls back to "torch" with a warning if "flashinfer" ' 'is requested but the package is not installed.', ) + group.add_argument( + '--use-same-sampling-seed-across-dp-ranks', + action='store_false', + dest='offset_sampling_seed_by_dp_rank', + default=True, + help='Use the same inference sampling seed on every data-parallel rank. ' + '--deterministic-mode also uses the same seed on every DP rank.', + ) group.add_argument( '--inference-dynamic-batching-async-sched-mode', type=str, @@ -2672,6 +2837,17 @@ def _add_inference_args(parser): 'By default, all graphs are limited by the decode limit of ' '`max_requests * (num_speculative_tokens + 1)`.', ) + group.add_argument( + '--inference-cuda-graph-max-tokens', + type=int, + default=512, + dest='inference_cuda_graph_max_tokens', + help='Token ceiling for the largest captured prefill/mixed CUDA ' + 'graph (default: 512). Clamped to at least the decode limit ' + '`max_requests * (num_speculative_tokens + 1)` and at most ' + '`max_tokens`. Ignored when --inference-cuda-graph-all-prefills ' + 'is set (which extends capture to the full `max_tokens`).', + ) group.add_argument( '--inference-dynamic-batching-cuda-graph-sizing-distribution', type=str, @@ -2782,6 +2958,10 @@ def _add_network_size_args(parser): "apply_dsa_kernel_fusion", "dsa_kernel_backend", "mamba_training_ssm_states_dtype", + # internal/derived: controlled only via --tensor-parallel-num-weight-shards + "gtp_weight_remat_size", + # internal/derived: controlled only via --expert-tensor-parallel-num-weight-shards + "expert_gtp_weight_remat_size", ] transformer_factory = ArgumentGroupFactory(TransformerConfig, exclude=exclude) transformer_group = transformer_factory.build_group(parser, "transformer configuration") @@ -3894,6 +4074,14 @@ def _add_checkpointing_args(parser): default=None, help='Do not load rng state when loading checkpoint.', ) + group.add_argument( + '--override-ckpt-iteration', + type=int, + default=None, + help='Override the iteration stored in the loaded checkpoint. ' + 'Also resets consumed_train_samples accordingly so the ' + 'data loader replays samples from that iteration onward.', + ) group.add_argument( '--use-dist-ckpt', action='store_true', @@ -4771,9 +4959,16 @@ def _add_moe_args(parser): '--moe-router-load-balancing-type', nargs='+', type=str, - choices=['aux_loss', 'seq_aux_loss', 'global_aux_loss', 'sinkhorn', 'none'], + choices=[ + 'aux_loss', + 'seq_aux_loss', + 'global_aux_loss', + 'sinkhorn', + 'quantile_balancing', + 'none', + ], default='aux_loss', - help='Determines the load balancing strategy for the router. "aux_loss" corresponds to the load balancing loss used in GShard and SwitchTransformer; "seq_aux_loss" corresponds to the load balancing loss used in DeepSeekV2, which computes the loss for each individual sample; "sinkhorn" corresponds to the balancing algorithm used in S-BASE, and "none" implies no load balancing. The default is "aux_loss".', + help='Determines the load balancing strategy for the router. "aux_loss" corresponds to the load balancing loss used in GShard and SwitchTransformer; "seq_aux_loss" corresponds to the load balancing loss used in DeepSeekV2, which computes the loss for each individual sample; "sinkhorn" corresponds to the balancing algorithm used in S-BASE; "quantile_balancing" (QB) uses dual coordinate descent on a per-expert bias to handle load balance internally; "none" implies no load balancing. The default is "aux_loss".', ) group.add_argument( '--moe-aux-loss-coeff', diff --git a/megatron/training/checkpointing.py b/megatron/training/checkpointing.py index bfa3564c6cd..0239fbea225 100644 --- a/megatron/training/checkpointing.py +++ b/megatron/training/checkpointing.py @@ -7,6 +7,7 @@ import multiprocessing import os import random +import re import shutil import sys import threading @@ -35,7 +36,7 @@ TorchDistSaveShardedStrategy, get_async_strategy, ) -from megatron.core.msc_utils import MultiStorageClientFeature, open_file +from megatron.core.msc_utils import maybe_msc from megatron.core.num_microbatches_calculator import update_num_microbatches from megatron.core.optimizer import DistributedOptimizer from megatron.core.rerun_state_machine import get_rerun_state_machine @@ -46,7 +47,7 @@ from .async_utils import get_save_and_finalize_callbacks, is_empty_async_queue, schedule_async_save from .global_vars import get_args from .one_logger_utils import on_save_checkpoint_start, on_save_checkpoint_success -from .utils import append_to_progress_log, is_last_rank, print_rank_0 +from .utils import append_to_progress_log, is_last_rank, print_rank_0, print_rank_last, warn_rank_0 try: from megatron.core.distributed.fsdp.src.megatron_fsdp.uneven_dtensor import ( @@ -140,12 +141,15 @@ def get_loaded_iteration(): return _LOADED_ITERATION -def check_checkpoint_args(checkpoint_args): +def check_checkpoint_args(checkpoint_args, skip_args: set[str] | None = None): """Ensure fixed arguments for a model are the same for the input arguments and the one retrieved from checkpoint.""" args = get_args() + skip_args = skip_args or set() def _compare(arg_name, old_arg_name=None, default=None): + if arg_name in skip_args: + return if old_arg_name is not None: ckpt_arg_name = old_arg_name else: @@ -182,22 +186,10 @@ def _compare(arg_name, old_arg_name=None, default=None): _compare('pipeline_model_parallel_size') -def isfile(filename) -> bool: - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - return msc.os.path.isfile(filename) - else: - return os.path.isfile(filename) - - def ensure_directory_exists(filename, check_parent=True): """Build filename's path if it does not already exists.""" dirname = os.path.dirname(filename) if check_parent else filename - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - msc.os.makedirs(dirname, exist_ok=True) - else: - os.makedirs(dirname, exist_ok=True) + maybe_msc.os.makedirs(dirname, exist_ok=True) def get_checkpoint_name( @@ -249,18 +241,18 @@ def get_checkpoint_name( return os.path.join(common_path, basename) -def get_load_checkpoint_path_by_args(args, load_arg="load"): +def get_load_checkpoint_path_by_args(args, load_arg='load'): """Get the checkpoint path based on the arguments.""" load_dir = getattr(args, load_arg) iteration, release = -1, False tracker_filename = 'because load directory is not defined' if load_dir is not None: tracker_filename = get_checkpoint_tracker_filename(load_dir) - if isfile(tracker_filename): + if maybe_msc.os.path.isfile(tracker_filename): iteration, release = read_metadata(tracker_filename) # Allow user to specify the loaded iteration. - if getattr(args, "ckpt_step", None): + if getattr(args, 'ckpt_step', None): iteration = args.ckpt_step return get_checkpoint_name(load_dir, iteration, release, return_base_dir=True) @@ -290,7 +282,7 @@ def find_checkpoint_rank_0(checkpoints_path, iteration, release=False): expert_parallel=False, expert_rank=0, ) - if isfile(filename): + if maybe_msc.os.path.isfile(filename): return filename # Look for checkpoint with no pipelining and expert parallelism @@ -304,7 +296,7 @@ def find_checkpoint_rank_0(checkpoints_path, iteration, release=False): expert_parallel=True, expert_rank=0, ) - if isfile(filename): + if maybe_msc.os.path.isfile(filename): return filename # Look for checkpoint with pipelining and no expert parallelism @@ -318,7 +310,7 @@ def find_checkpoint_rank_0(checkpoints_path, iteration, release=False): expert_parallel=False, expert_rank=0, ) - if isfile(filename): + if maybe_msc.os.path.isfile(filename): return filename # Look for checkpoint with pipelining and expert parallelism @@ -332,7 +324,7 @@ def find_checkpoint_rank_0(checkpoints_path, iteration, release=False): expert_parallel=True, expert_rank=0, ) - if isfile(filename): + if maybe_msc.os.path.isfile(filename): return filename # Look for a distributed checkpoint @@ -355,7 +347,7 @@ def checkpoint_exists(checkpoints_path): if checkpoints_path is None: return False path = get_checkpoint_tracker_filename(checkpoints_path) - return isfile(path) + return maybe_msc.os.path.isfile(path) def read_metadata(tracker_filename): @@ -364,7 +356,7 @@ def read_metadata(tracker_filename): iteration = -1 release = False - with open_file(tracker_filename, 'r') as f: + with maybe_msc.open(tracker_filename, 'r') as f: metastring = f.read().strip() try: iteration = int(metastring) @@ -403,6 +395,24 @@ def read_metadata(tracker_filename): return max_iter, release +def read_frozen_resume_iteration(save_dir): + """Resume iteration for a ``--freeze-all-layers`` run. + + Returns the integer recorded in ``save_dir``'s progress tracker + (``latest_checkpointed_iteration.txt``), or ``0`` when there is none -- i.e. a + fresh run, equivalent to ``--finetune`` on the first launch. Unlike + :func:`read_metadata` (which returns ``-1`` / errors on a missing or malformed + file), a missing tracker here simply means "start from the beginning". + """ + if save_dir is None: + return 0 + tracker_filename = get_checkpoint_tracker_filename(save_dir) + if not maybe_msc.os.path.isfile(tracker_filename): + return 0 + iteration, _release = read_metadata(tracker_filename) + return iteration + + def get_rng_state( ckpt_format: str, tp_group: torch.distributed.ProcessGroup, @@ -500,10 +510,10 @@ def _build_sharded_state_dict_metadata( args, 'use_layer_wise_distributed_optimizer', False ) - if has_distributed_optimizer and args.ckpt_format == "fsdp_dtensor": + if has_distributed_optimizer and args.ckpt_format == 'fsdp_dtensor': metadata['distrib_optim_sharding_type'] = 'fsdp_dtensor' - if has_distributed_optimizer and args.ckpt_format != "fsdp_dtensor": + if has_distributed_optimizer and args.ckpt_format != 'fsdp_dtensor': if args.dist_ckpt_optim_fully_reshardable: metadata['distrib_optim_sharding_type'] = 'fully_reshardable' metadata['distrib_optim_fully_reshardable_mem_efficient'] = ( @@ -540,16 +550,16 @@ def save_grads(save_dir, state_dict, iteration, grad_label): tp_rank = mpu.get_tensor_model_parallel_rank() assert save_dir is not None assert iteration is not None - save_dir = os.path.join(save_dir, grad_label, f"iter_{iteration:07d}") + save_dir = os.path.join(save_dir, grad_label, f'iter_{iteration:07d}') os.makedirs(save_dir, exist_ok=True) # Save state_dict. - checkpoint_name = f"mp_rank_{tp_rank:02d}" + checkpoint_name = f'mp_rank_{tp_rank:02d}' if mpu.get_pipeline_model_parallel_world_size() > 1: - checkpoint_name += f"_{pp_rank:03d}" + checkpoint_name += f'_{pp_rank:03d}' if mpu.get_expert_model_parallel_world_size() > 1: - checkpoint_name += f"_{ep_rank:03d}" - full_save_path = os.path.join(save_dir, f"{checkpoint_name}.pth") + checkpoint_name += f'_{ep_rank:03d}' + full_save_path = os.path.join(save_dir, f'{checkpoint_name}.pth') # Convert back to dict (e.g., from collections.defaultdict) for easy loading later. torch.save(dict(state_dict), full_save_path) @@ -618,6 +628,10 @@ def save_checkpoint( # Only rank zero of the data parallel writes to the disk. model = unwrap_model(model) + # --freeze-all-layers: weights are frozen, so skip the weight write and its finalizers; dump + # progress is recorded separately via the post-logits finalize (see the async block below). + skip_weight_ckpt = getattr(args, 'freeze_all_layers', False) + # Handle non_persistent_ckpt flag. Besides overwriting `args.save` and # `args.use_dist_ckpt`, non-persistent global ckpt requires no additional logic ckpt_type = CheckpointType.GLOBAL if args.use_dist_ckpt else CheckpointType.LEGACY @@ -681,9 +695,14 @@ def save_checkpoint( return_base_dir=return_base_dir, ) - # Save dataloader state if the dataloader supports it (currently only Megatron Energon). + # Save dataloader state if the external dataloader supports it. maybe_save_dataloader_state( - train_data_iterator, iteration, getattr(args, "dataloader_save", None) + train_data_iterator, + iteration, + getattr(args, 'dataloader_save', None), + tp_group=tp_group, + pp_group=pp_group, + dp_group=dp_group, ) # Save distributed optimizer's custom parameter state. @@ -699,7 +718,7 @@ def save_checkpoint( optimizer.save_parameter_state(optim_checkpoint_name) # LayerWiseDistributedOptimizer save optimizer state to file on different ranks - if getattr(args, "use_layer_wise_distributed_optimizer", False) and args.ckpt_format == 'torch': + if getattr(args, 'use_layer_wise_distributed_optimizer', False) and args.ckpt_format == 'torch': dp_rank = mpu.get_data_parallel_rank() optim_checkpoint_name = os.path.join( os.path.dirname(checkpoint_name), f"layer_wise_optimizer_{dp_rank}.pt" @@ -739,7 +758,7 @@ def save_checkpoint( # exactly one rank. Neither dp_rank==0 nor edp_rank==0 alone covers all shards when # the dense and expert parallelism layouts disagree (e.g. TP > EP*ETP); the union # does, with at most one rank per (tp_rank, ep_rank) inside any DP group. - if ( + if not skip_weight_ckpt and ( not torch.distributed.is_initialized() or ckpt_type != CheckpointType.LEGACY or dp_rank == 0 @@ -767,7 +786,7 @@ def save_checkpoint( ) state_dict['num_floating_point_operations_so_far'] = num_floating_point_operations_so_far - if ckpt_type == CheckpointType.GLOBAL and ckpt_format == "torch_dist": + if ckpt_type == CheckpointType.GLOBAL and ckpt_format == 'torch_dist': if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0: # TODO Handle non-empty directories (e.g., after a crash during saving). ensure_directory_exists(checkpoint_name, check_parent=False) @@ -799,7 +818,7 @@ def save_checkpoint( checkpointing_context['load_strategy'], 'cached_global_metadata', None ) if cached_global_metadata is not None: - logger.debug("Plugging in the read metadata from the load strategy...") + logger.debug('Plugging in the read metadata from the load strategy...') save_strategy.cached_global_metadata = cached_global_metadata else: logger.debug( @@ -843,33 +862,33 @@ def save_checkpoint( # [ModelOpt]: save sharded modelopt_state if has_nvidia_modelopt: save_sharded_modelopt_state(model, checkpoint_name, (args.ckpt_format, 1)) - elif ckpt_type == CheckpointType.GLOBAL and ckpt_format in ["torch_dcp", "fsdp_dtensor"]: + elif ckpt_type == CheckpointType.GLOBAL and ckpt_format in ['torch_dcp', 'fsdp_dtensor']: if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0: # TODO Handle non-empty directories (e.g., after a crash during saving). ensure_directory_exists(checkpoint_name, check_parent=False) - if ckpt_format == "fsdp_dtensor": + if ckpt_format == 'fsdp_dtensor': state_dict = preprocess_fsdp_dtensor_state_dict(args, state_dict, model[0]) if args.async_save: planner = torch.distributed.checkpoint.DefaultSavePlanner() coordinator_rank = 0 _, async_modules = get_async_strategy(args.async_strategy) - FileSystemWriterAsync = async_modules["FileSystemWriterAsync"] - save_state_dict_async_plan = async_modules["save_state_dict_async_plan"] + FileSystemWriterAsync = async_modules['FileSystemWriterAsync'] + save_state_dict_async_plan = async_modules['save_state_dict_async_plan'] _cpu_shm = getattr(args, 'async_ckpt_use_cpu_shm', False) _writer_kwargs = {} if _cpu_shm: if ( - "use_cpu_shm_for_gpu_tensors" + 'use_cpu_shm_for_gpu_tensors' in inspect.signature(FileSystemWriterAsync.__init__).parameters ): - _writer_kwargs["use_cpu_shm_for_gpu_tensors"] = True + _writer_kwargs['use_cpu_shm_for_gpu_tensors'] = True else: raise AssertionError( - "Installed nvidia-resiliency-ext does not support " - "use_cpu_shm_for_gpu_tensors. Update nvidia-resiliency-ext " - "to use --async-ckpt-use-cpu-shm." + 'Installed nvidia-resiliency-ext does not support ' + 'use_cpu_shm_for_gpu_tensors. Update nvidia-resiliency-ext ' + 'to use --async-ckpt-use-cpu-shm.' ) fs_storage_writer = FileSystemWriterAsync( checkpoint_name, @@ -886,9 +905,6 @@ def save_checkpoint( planner=planner, enable_cache=args.ckpt_assume_constant_structure, ) - async_save_request = get_save_and_finalize_callbacks( - fs_storage_writer, save_state_dict_ret - ) async_save_request = get_save_and_finalize_callbacks( fs_storage_writer, save_state_dict_ret, args.async_strategy ) @@ -963,7 +979,9 @@ def save_checkpoint( torch.distributed.barrier() # And update the latest iteration - if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0: + if not skip_weight_ckpt and ( + not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0 + ): tracker_filename = get_checkpoint_tracker_filename(save_dir) if ckpt_type == CheckpointType.LOCAL: @@ -1003,6 +1021,8 @@ def _rank_and_size(explicit_rank, group, mpu_rank_fn, mpu_size_fn): mpu.get_pipeline_model_parallel_rank, mpu.get_pipeline_model_parallel_world_size, ) + gtp_remat_rank = mpu.get_gtp_weight_remat_rank() + 1 + gtp_remat_size_to_print = mpu.get_gtp_weight_remat_world_size() def iter_finalize_fn(): prev_iteration = 0 @@ -1010,18 +1030,17 @@ def iter_finalize_fn(): args, 'save_retain_interval', None ) # For backwards compatibility of tests. if save_retain_interval is not None: - if os.path.exists( - tracker_filename - ): # TODO: Make this work with MSC remote paths? - with open_file(tracker_filename, 'r') as f: + if maybe_msc.os.path.exists(tracker_filename): + with maybe_msc.open(tracker_filename, 'r') as f: prev_iteration = int(f.read().strip()) - with open_file(tracker_filename, 'w') as f: - f.write("release" if release else str(iteration)) + with maybe_msc.open(tracker_filename, 'w') as f: + f.write('release' if release else str(iteration)) print_rank_0( - f" [{datetime.now().strftime('%Y-%m-%d %H:%M:%S.%f')}] successfully saved " - f"checkpoint from iteration {int(iteration):7d} to {args.save} " - f"[ t {tensor_mp_rank}/{tp_size_to_print}, " - f"p {pipeline_mp_rank}/{pp_size_to_print} ]" + f' [{datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f")}] successfully saved ' + f'checkpoint from iteration {int(iteration):7d} to {args.save} ' + f'[ t {tensor_mp_rank}/{tp_size_to_print}, ' + f'gtp_remat {gtp_remat_rank}/{gtp_remat_size_to_print}, ' + f'p {pipeline_mp_rank}/{pp_size_to_print} ]' ) if args.log_progress and args.async_save: append_to_progress_log( @@ -1081,7 +1100,7 @@ def iter_finalize_fn(): iter_finalize_fn() # Additional callback for one_logger (last rank) - if not torch.distributed.is_initialized() or is_last_rank(): + if not skip_weight_ckpt and (not torch.distributed.is_initialized() or is_last_rank()): def onelogger_finalize_fn(): on_save_checkpoint_success(productive_metrics, args.async_save) @@ -1093,7 +1112,7 @@ def onelogger_finalize_fn(): onelogger_finalize_fn() # Additional callback for wandb (last rank) - if not torch.distributed.is_initialized() or is_last_rank(): + if not skip_weight_ckpt and (not torch.distributed.is_initialized() or is_last_rank()): def wandb_finalize_fn(): wandb_utils.on_save_checkpoint_success( @@ -1107,14 +1126,56 @@ def wandb_finalize_fn(): wandb_finalize_fn() if args.async_save: - schedule_async_save(async_save_request) + # Schedule logits flush AFTER the checkpoint request so the persistent + # worker processes checkpoint preload first (unblocking the main + # thread), then writes logits in the background. Finalize_fns are + # moved from the checkpoint request to the logits request so that + # "success" callbacks only fire after both writes are confirmed. + from megatron.training.distillation import get_logits_saver + + logits_saver = get_logits_saver() + if logits_saver is not None: + # In frozen-dump mode there is no checkpoint request (async_save_request is None); the + # logits request then owns its own finalize_fns. + if async_save_request is not None: + logits_finalize_fns = async_save_request.finalize_fns.copy() + async_save_request.finalize_fns.clear() + else: + logits_finalize_fns = [] + # Record run progress AFTER the logits tar is confirmed written, so a resumed job + # never skips or replays a window. Written by a single rank -- the last rank, which + # lives on the last pipeline stage where the logits saver is attached (get_logits_saver + # is None on earlier stages, including global rank 0 when PP > 1). + if skip_weight_ckpt and (not torch.distributed.is_initialized() or is_last_rank()): + + def progress_finalize_fn(): + tracker_filename = get_checkpoint_tracker_filename(args.save) + with maybe_msc.open(tracker_filename, 'w') as f: + f.write(str(iteration)) + print_rank_last( + f" recorded logits-dump progress: iteration " + f"{iteration} to {tracker_filename}" + ) + + logits_finalize_fns.append(progress_finalize_fn) + async_request_cls = get_async_strategy(args.async_strategy)[1]['AsyncRequest'] + async_logits_request = async_request_cls( + async_fn=logits_saver._write_batched_tar, + async_fn_args=logits_saver.take_pending_data(), + finalize_fns=logits_finalize_fns, + ) + + if async_save_request is not None: + schedule_async_save(async_save_request) + if logits_saver is not None: + schedule_async_save(async_logits_request) print_rank_0( - f" [{datetime.now().strftime('%Y-%m-%d %H:%M:%S.%f')}] scheduled " - f"an async checkpoint save at iteration {iteration:7d} to {save_dir}" + f' [{datetime.now().strftime("%Y-%m-%d %H:%M:%S.%f")}] scheduled ' + f'an async checkpoint save at iteration {iteration:7d} to {save_dir}' ) end_misc = time() - logger.debug(f"rank: {rank}, takes {end_misc - start_misc} to finalize ckpt save ") + logger.debug(f'rank: {rank}, takes {end_misc - start_misc} to finalize ckpt save ') if not args.async_save: # Add a barrier so that all ranks wait for finalization to complete @@ -1182,7 +1243,7 @@ def cleanup_old_non_persistent_checkpoint(save_dir, leave_ckpt_num=1, do_async=F return save_dir = Path(save_dir) - iter_prefix = "iter_" + iter_prefix = 'iter_' iter_ckpts = save_dir.rglob(f'{iter_prefix}*') sorted_iter_ckpts = sorted( iter_ckpts, key=lambda ckpt_name: int(ckpt_name.name[len(iter_prefix) :]) @@ -1203,12 +1264,13 @@ def remove_iter_ckpts(_iter_ckpts): remove_iter_ckpts(rm_iter_ckpts) -def maybe_save_dataloader_state(train_iterator, iteration, dataloader_save_path): +def maybe_save_dataloader_state( + train_iterator, iteration, dataloader_save_path, *, tp_group=None, pp_group=None, dp_group=None +): """Saves dataloader state if the dataloader supports it. - Currently, this is only used by Megatron Energon dataloader (multimodal) to store its state at a - specific iteration. The Megatron built-in dataloader (text-only) creates index files upfront - to track its state. + External dataloaders use this to store state at a specific iteration. The Megatron built-in + dataloader creates index files upfront to track its state. If the provided dataloader has `save_state` method, then it is called to save the state. Otherwise, no state is saved. @@ -1217,9 +1279,12 @@ def maybe_save_dataloader_state(train_iterator, iteration, dataloader_save_path) train_iterator (iterable): Train dataloader. iteration (int): Current iteration. dataloader_save_path (str): Path where the dataloader state is saved. + tp_group (ProcessGroup): Tensor-parallel group, or MPU fallback when unset. + pp_group (ProcessGroup): Pipeline-parallel group, or MPU fallback when unset. + dp_group (ProcessGroup): Data-parallel group, or MPU fallback when unset. """ # If no dataloader or saving path is provided, exit early, otherwise, raise an error. - if train_iterator is None or dataloader_save_path is None or dataloader_save_path == "": + if train_iterator is None or dataloader_save_path is None or dataloader_save_path == '': return # If dataloader doesn't support saving state, raise an error. @@ -1230,26 +1295,47 @@ def maybe_save_dataloader_state(train_iterator, iteration, dataloader_save_path) # Save dataloader state for each data parallel rank only once. first_rank = ( - mpu.is_pipeline_first_stage(ignore_virtual=True) - and mpu.get_tensor_model_parallel_rank() == 0 + get_pg_rank(pp_group) == 0 + if pp_group is not None + else mpu.is_pipeline_first_stage(ignore_virtual=True) + ) and ( + get_pg_rank(tp_group) == 0 + if tp_group is not None + else mpu.get_tensor_model_parallel_rank() == 0 ) if not first_rank: return - dp_rank = mpu.get_data_parallel_rank() - if dp_rank == 0: - print(f"saving dataloader checkpoint at iteration {iteration} to {dataloader_save_path}") + dp_rank = get_pg_rank(dp_group) if dp_group is not None else mpu.get_data_parallel_rank() train_dataloader_state_dict = train_iterator.iterable.save_state() + if dp_rank == 0: + print(f'saving dataloader checkpoint at iteration {iteration} to {dataloader_save_path}') data_state_save_path = get_checkpoint_name( - dataloader_save_path, iteration, basename=f'train_dataloader_dprank{dp_rank:03d}.pt' + dataloader_save_path, + iteration, + pipeline_parallel=( + get_pg_size(pp_group) > 1 + if pp_group is not None + else mpu.get_pipeline_model_parallel_world_size() > 1 + ), + # Dataloader state is sharded only by DP rank. Keep it in the canonical TP0/PP0 directory. + tensor_rank=0, + pipeline_rank=0, + expert_parallel=False, + expert_rank=0, + basename=f'train_dataloader_dprank{dp_rank:03d}.pt', ) - torch.distributed.barrier(group=mpu.get_data_parallel_group()) + data_parallel_group = dp_group if dp_group is not None else mpu.get_data_parallel_group() + torch.distributed.barrier(group=data_parallel_group) - if mpu.get_data_parallel_rank() == 0: + if dp_rank == 0: ensure_directory_exists(data_state_save_path) - torch.distributed.barrier(group=mpu.get_data_parallel_group()) + torch.distributed.barrier(group=data_parallel_group) + + if train_dataloader_state_dict is None: + return dataloader_save_dict = {} dataloader_save_dict['dataloader_state_dict'] = train_dataloader_state_dict @@ -1277,11 +1363,11 @@ def generate_state_dict( state_dict['iteration'] = iteration for i in range(len(model)): - key = "model" + key = 'model' if len(model) > 1: - key = f"model{i}" + key = f'model{i}' - if args.ckpt_format == "torch_dist": + if args.ckpt_format == 'torch_dist': model_sd = model[i].sharded_state_dict( **( model_sd_kwargs @@ -1300,8 +1386,7 @@ def generate_state_dict( # Optimizer stuff. if not args.no_save_optim: if optimizer is not None and not optimizer.is_stub_optimizer: - - if args.ckpt_format == "torch_dist": + if args.ckpt_format == 'torch_dist': optimizer_sd = optimizer.sharded_state_dict( state_dict, **( @@ -1315,11 +1400,11 @@ def generate_state_dict( } ), ) - elif args.ckpt_format == "fsdp_dtensor": + elif args.ckpt_format == 'fsdp_dtensor': if optim_sd_kwargs is None: optim_sd_kwargs = {} - if "metadata" not in optim_sd_kwargs: - optim_sd_kwargs["metadata"] = {} + if 'metadata' not in optim_sd_kwargs: + optim_sd_kwargs['metadata'] = {} optim_sd_kwargs['metadata'].update(_build_sharded_state_dict_metadata(args)) optimizer_sd = optimizer.sharded_state_dict(state_dict, **optim_sd_kwargs) else: @@ -1336,21 +1421,21 @@ def generate_state_dict( # RNG states. if not args.no_save_rng and rng_state: - state_dict["rng_state"] = rng_state + state_dict['rng_state'] = rng_state return state_dict def preprocess_fsdp_dtensor_state_dict(args, raw_state_dict, model): state_dict = raw_state_dict.copy() - handle_fp8_extra_state_case(state_dict["model"]) + handle_fp8_extra_state_case(state_dict['model']) if args.swiglu: - if "optimizer" in state_dict: + if 'optimizer' in state_dict: model_state_dict, optimizer_state_dict = handle_swiglu_in_state_dict( - model, state_dict["model"], state_dict["optimizer"] + model, state_dict['model'], state_dict['optimizer'] ) - state_dict["model"] = model_state_dict - state_dict["optimizer"] = optimizer_state_dict + state_dict['model'] = model_state_dict + state_dict['optimizer'] = optimizer_state_dict else: model_state_dict, _ = handle_swiglu_in_state_dict(model, state_dict["model"], None) state_dict["model"] = model_state_dict @@ -1364,7 +1449,7 @@ def preprocess_fsdp_dtensor_state_dict(args, raw_state_dict, model): model_state_dict, _ = handle_gdn_in_state_dict(model, state_dict["model"], None) state_dict["model"] = model_state_dict if args.num_experts: - state_dict["model"] = handle_experts_in_state_dict(state_dict["model"], args.num_experts) + state_dict['model'] = handle_experts_in_state_dict(state_dict['model'], args.num_experts) preprocess_state_dict_for_uneven_dtensor(state_dict) return state_dict @@ -1427,7 +1512,7 @@ def fix_query_key_value_ordering(model, checkpoint_version): elif checkpoint_version == 1.0: fixed_param = _transpose_first_dim(param.data, 3, False, model) else: - print_rank_0(f"Invalid checkpoint version {checkpoint_version}.") + print_rank_0(f'Invalid checkpoint version {checkpoint_version}.') sys.exit() param.data.copy_(fixed_param) if name.endswith(('.key_value.weight', '.key_value.bias')): @@ -1436,7 +1521,7 @@ def fix_query_key_value_ordering(model, checkpoint_version): elif checkpoint_version == 1.0: fixed_param = _transpose_first_dim(param.data, 2, False, model) else: - print_rank_0(f"Invalid checkpoint version {checkpoint_version}.") + print_rank_0(f'Invalid checkpoint version {checkpoint_version}.') sys.exit() param.data.copy_(fixed_param) print_rank_0( @@ -1448,18 +1533,18 @@ def fix_query_key_value_ordering(model, checkpoint_version): def _get_non_persistent_iteration(non_persistent_global_dir, args, checkpointing_context=None): if args.non_persistent_ckpt_type is None: return -1 - elif args.non_persistent_ckpt_type == "global": + elif args.non_persistent_ckpt_type == 'global': tracker_filename = get_checkpoint_tracker_filename(non_persistent_global_dir) - if isfile(tracker_filename): + if maybe_msc.os.path.isfile(tracker_filename): iteration, release = read_metadata(tracker_filename) if release: - raise RuntimeError('Non-persistent checkpoint can\'t be a release checkpoint') + raise RuntimeError("Non-persistent checkpoint can't be a release checkpoint") else: iteration = -1 print_rank_0('WARNING: could not find the metadata file {}'.format(tracker_filename)) print_rank_0(' will not load any non-persistent checkpoint') return iteration - elif args.non_persistent_ckpt_type == "local": + elif args.non_persistent_ckpt_type == 'local': return checkpointing_context['local_checkpoint_manager'].find_latest() else: assert False, ( @@ -1482,7 +1567,7 @@ def _load_non_persistent_base_checkpoint( Depending on the non_persistent_ckpt_type, different logic may be required. """ assert args.non_persistent_ckpt_type is not None - if args.non_persistent_ckpt_type == "global": + if args.non_persistent_ckpt_type == 'global': if not rank0: print_rank_0( f'Loading from a non-persistent checkpoint (non-persistent iter {non_persistent_iteration})' @@ -1498,7 +1583,7 @@ def _load_non_persistent_base_checkpoint( dp_cp_group=dp_cp_group, expt_dp_group=expt_dp_group, ) - elif args.non_persistent_ckpt_type == "local": + elif args.non_persistent_ckpt_type == 'local': intermediate_state_dict, checkpoint_name = checkpointing_context[ 'local_checkpoint_manager' ].load() @@ -1545,7 +1630,10 @@ def _load_global_dist_base_checkpoint( ) checkpoint_name = get_checkpoint_name(load_dir, iteration, release, return_base_dir=True) - load_strategy = TorchDistLoadShardedStrategy(cache_metadata=args.ckpt_assume_constant_structure) + load_strategy = TorchDistLoadShardedStrategy( + cache_metadata=args.ckpt_assume_constant_structure, + stream_ckpt_dequant=args.stream_ckpt_dequant, + ) # NOTE: `args.ckpt_fully_parallel_load` applies to both persistent and non-persistent checkpoints. if args.ckpt_fully_parallel_load: if args.ckpt_fully_parallel_load_process_group == 'dp': @@ -1567,7 +1655,7 @@ def _load_global_dist_base_checkpoint( load_strategy, process_group, exchange_algo=args.ckpt_fully_parallel_load_exchange_algo ) if checkpointing_context is not None: - checkpointing_context["load_strategy"] = load_strategy + checkpointing_context['load_strategy'] = load_strategy state_dict = dist_checkpointing.load( sharded_state_dict, checkpoint_name, @@ -1581,26 +1669,21 @@ def _load_global_dist_base_checkpoint( def _get_checkpoint_format(checkpoint_name, args): """Get the format of an existing checkpoint.""" - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - checkpoint_dir = msc.Path(checkpoint_name) - is_torch_ckpt = any([f.name.startswith("mp_rank_0") for f in checkpoint_dir.iterdir()]) - is_torch_dcp = checkpoint_dir.joinpath(".metadata").exists() - else: - is_torch_ckpt = any([f.startswith("mp_rank_0") for f in os.listdir(checkpoint_name)]) - is_torch_dcp = os.path.exists(os.path.join(checkpoint_name, ".metadata")) + checkpoint_dir = maybe_msc.Path(checkpoint_name) + is_torch_ckpt = any([f.name.startswith('mp_rank_0') for f in checkpoint_dir.iterdir()]) + is_torch_dcp = checkpoint_dir.joinpath('.metadata').exists() ckpt_format = None if dist_checkpointing.check_is_distributed_checkpoint(checkpoint_name): - ckpt_format = "torch_dist" + ckpt_format = 'torch_dist' elif is_torch_ckpt: - ckpt_format = "torch" + ckpt_format = 'torch' elif is_torch_dcp: - ckpt_format = "torch_dcp" - if getattr(args, "use_megatron_fsdp", False): - ckpt_format = "fsdp_dtensor" + ckpt_format = 'torch_dcp' + if getattr(args, 'use_megatron_fsdp', False): + ckpt_format = 'fsdp_dtensor' else: - raise NotImplementedError(f"unknown checkpoint format in {checkpoint_name}") + raise NotImplementedError(f'unknown checkpoint format in {checkpoint_name}') return ckpt_format @@ -1613,6 +1696,7 @@ def _load_base_checkpoint( checkpointing_context=None, dp_cp_group=None, expt_dp_group=None, + gpt_compat_layer_maps=None, ): """Load the base state_dict from the given directory @@ -1631,11 +1715,11 @@ def _load_base_checkpoint( tracker_filename = 'because load directory is not defined' if load_dir is not None: tracker_filename = get_checkpoint_tracker_filename(load_dir) - if isfile(tracker_filename): + if maybe_msc.os.path.isfile(tracker_filename): iteration, release = read_metadata(tracker_filename) # Allow user to specify the loaded iteration. - if getattr(args, "ckpt_step", None): + if getattr(args, 'ckpt_step', None): iteration = args.ckpt_step # Record the iteration loaded (stored separately from args to avoid @@ -1670,14 +1754,14 @@ def _load_base_checkpoint( torch.distributed.barrier() sys.exit() - return None, "", False, None + return None, '', False, None # Determine the type of the checkpoint on disk. checkpoint_name = get_checkpoint_name(load_dir, iteration, release, return_base_dir=True) ckpt_format = _get_checkpoint_format(checkpoint_name, args) if not rank0: - dist_infix = "distributed " if ckpt_format == "torch_dist" else "" + dist_infix = 'distributed ' if ckpt_format == 'torch_dist' else '' if release: print_rank_0(f' loading release {dist_infix}checkpoint from {load_dir}') else: @@ -1688,7 +1772,7 @@ def _load_base_checkpoint( ckpt_type = None # Handle global distributed checkpoint - if ckpt_format == "torch_dist": + if ckpt_format == 'torch_dist': return _load_global_dist_base_checkpoint( load_dir, args, @@ -1700,7 +1784,7 @@ def _load_base_checkpoint( dp_cp_group=dp_cp_group, expt_dp_group=expt_dp_group, ) - elif ckpt_format == "torch": + elif ckpt_format == 'torch': ckpt_type = CheckpointType.LEGACY # Handle global legacy checkpoint if rank0: @@ -1732,7 +1816,7 @@ def _load_base_checkpoint( print('could not load the checkpoint') print(e) sys.exit() - elif ckpt_format == "torch_dcp": + elif ckpt_format == 'torch_dcp': ckpt_type = CheckpointType.TORCH_DCP if rank0: @@ -1749,10 +1833,12 @@ def _load_base_checkpoint( torch.distributed.checkpoint.load_state_dict( state_dict=state_dict, storage_reader=fs_storage_reader ) - elif ckpt_format == "fsdp_dtensor": - assert HAVE_MEGATRON_FSDP, "Should not be called if Megatron-FSDP is not available." + elif ckpt_format == 'fsdp_dtensor': + assert HAVE_MEGATRON_FSDP, 'Should not be called if Megatron-FSDP is not available.' if rank0: - return {}, checkpoint_name, release, CheckpointType.FSDP_DTENSOR + state_dict = {'args': None, 'iteration': None, 'checkpoint_version': None} + torch.distributed.checkpoint.load(state_dict=state_dict, checkpoint_id=checkpoint_name) + return state_dict, checkpoint_name, release, CheckpointType.FSDP_DTENSOR state_dict = sharded_state_dict raw_optimizer_state_dict = ( @@ -1761,9 +1847,28 @@ def _load_base_checkpoint( raw_model_state_dict = state_dict["model"].copy() if "model" in state_dict else None model = state_dict.pop("_model") state_dict = preprocess_fsdp_dtensor_state_dict(args, state_dict, model[0]) + fs_storage_reader = torch.distributed.checkpoint.FileSystemReader(checkpoint_name) + state_dict_metadata = fs_storage_reader.read_metadata().state_dict_metadata + if gpt_compat_layer_maps is not None: + from megatron.core.dist_checkpointing.gpt_checkpoint_interop import ( + retarget_fsdp_state_dict_to_gpt_checkpoint, + ) + + state_dict['model'] = retarget_fsdp_state_dict_to_gpt_checkpoint( + state_dict['model'], + gpt_compat_layer_maps, + tuple(key for key in state_dict_metadata if key.startswith('model.')), + checkpoint_prefix='model', + ) + if 'optimizer' in state_dict: + state_dict['optimizer'] = retarget_fsdp_state_dict_to_gpt_checkpoint( + state_dict['optimizer'], + gpt_compat_layer_maps, + tuple(key for key in state_dict_metadata if key.startswith('optimizer.')), + checkpoint_prefix='optimizer', + ) ckpt_type = CheckpointType.FSDP_DTENSOR - fs_storage_reader = torch.distributed.checkpoint.FileSystemReader(checkpoint_name) allow_partial_load = not getattr(args, 'strict_fsdp_dtensor_load', False) # Stash state_dict keys that have no presence in the checkpoint. @@ -1800,12 +1905,12 @@ def _load_base_checkpoint( state_dict.update(_stashed_keys) if raw_optimizer_state_dict is not None: - state_dict["optimizer"] = raw_optimizer_state_dict + state_dict['optimizer'] = raw_optimizer_state_dict if raw_model_state_dict is not None: - state_dict["model"] = raw_model_state_dict + state_dict['model'] = raw_model_state_dict else: - raise NotImplementedError(f"checkpoint format {ckpt_format} not supported") + raise NotImplementedError(f'checkpoint format {ckpt_format} not supported') return state_dict, checkpoint_name, release, ckpt_type @@ -1877,10 +1982,10 @@ def _set_arg(arg_name, old_arg_name=None, force=False): checkpoint_value = getattr(checkpoint_args, arg_name, None) if checkpoint_value is not None: - print_rank_0(f"Setting {arg_name} to {checkpoint_value} from checkpoint") + print_rank_0(f'Setting {arg_name} to {checkpoint_value} from checkpoint') setattr(args, arg_name, checkpoint_value) else: - print_rank_0(f"Checkpoint did not provide arguments {arg_name}") + print_rank_0(f'Checkpoint did not provide arguments {arg_name}') # Model args. _set_arg('num_layers') @@ -1975,6 +2080,86 @@ def _set_arg(arg_name, old_arg_name=None, force=False): return args, checkpoint_args +def _maybe_setup_gpt_to_hybrid_load(args, ckpt_args, model): + """Detect a GPT (pure transformer) checkpoint being loaded into a HybridModel run. + + Returns ``(layer_maps, load_optim)`` where ``layer_maps`` is used to retarget + the run's sharded state dict at the GPT checkpoint's keys (see + ``megatron.core.dist_checkpointing.gpt_checkpoint_interop``) and ``load_optim`` is + True when the GPT optimizer state should be translated and loaded as well. + Returns ``(None, False)`` when checkpoint and runtime already agree. Raises + RuntimeError for combinations that cannot be loaded. + """ + from megatron.core.dist_checkpointing.gpt_checkpoint_interop import gpt_compatible_layer_maps + from megatron.core.models.hybrid.hybrid_model import HybridModel + + def _contains_hybrid_model(module): + # Megatron-FSDP and Float16Module both retain the wrapped module under + # ``module`` but are intentionally not handled by the regular + # ``unwrap_model`` helper. + while module is not None: + if isinstance(module, HybridModel): + return True + module = getattr(module, 'module', None) + return False + + runtime_is_hybrid = any(_contains_hybrid_model(m) for m in model) + ckpt_pattern = getattr(ckpt_args, 'hybrid_layer_pattern', None) or getattr( + ckpt_args, 'hybrid_override_pattern', None + ) + if runtime_is_hybrid == bool(ckpt_pattern): + return None, False + if not runtime_is_hybrid: + raise RuntimeError( + f'The checkpoint was saved by a hybrid model run (hybrid layer pattern ' + f'{ckpt_pattern!r}) but the current run builds a non-hybrid model. Load it ' + f'with the hybrid training entrypoint instead.' + ) + + # GPT checkpoint feeding a HybridModel run: translate the sharded state + # dict at load time instead of converting the checkpoint on disk. + if not args.hybrid_layer_pattern: + raise RuntimeError( + 'Loading a GPT checkpoint into a hybrid model requires ' + '--hybrid-layer-pattern so checkpoint layers can be paired with ' + 'hybrid layer positions.' + ) + try: + layer_maps = gpt_compatible_layer_maps(args.hybrid_layer_pattern) + except ValueError as exc: + raise RuntimeError(f'Cannot load a GPT checkpoint into this hybrid model: {exc}') from exc + + ckpt_num_layers = getattr(ckpt_args, 'num_layers', None) + if ckpt_num_layers is not None and ckpt_num_layers != layer_maps.num_gpt_layers: + raise RuntimeError( + f'Hybrid layer pattern {args.hybrid_layer_pattern!r} pairs with a GPT ' + f'checkpoint of {layer_maps.num_gpt_layers} layers, but the checkpoint ' + f'has num_layers={ckpt_num_layers}.' + ) + + # The optimizer state is loaded unless the user opts out or the GPT run saved + # no optimizer state. Fresh layers (e.g. Mamba) have no counterpart in the GPT + # checkpoint, so their optimizer state stays freshly initialized; warn about it. + load_optim = not args.no_load_optim and not getattr(ckpt_args, 'no_save_optim', False) + if load_optim and layer_maps.fresh_init: + print_rank_0( + f'> WARNING: {len(layer_maps.fresh_init)} hybrid layer(s) have no GPT ' + f'counterpart (e.g. Mamba positions); their weights and optimizer state ' + f'start from a fresh initialization while the GPT-sourced attention and ' + f'MLP layers load their optimizer state. Pass --no-load-optim to start ' + f"every layer's optimizer state fresh." + ) + + print_rank_0( + f'> loading a GPT checkpoint into the hybrid model: {layer_maps.num_gpt_layers} ' + f'GPT layers feed {len(layer_maps.attention_to_gpt)} attention and ' + f'{len(layer_maps.mlp_to_gpt)} MLP positions; {len(layer_maps.fresh_init)} ' + f'hybrid layers keep their fresh initialization' + + ('; optimizer state will be loaded.' if load_optim else '; optimizer state starts fresh.') + ) + return layer_maps, load_optim + + def load_checkpoint( ddp_model, optimizer, @@ -2004,6 +2189,22 @@ def load_checkpoint( args = get_args() load_dir = getattr(args, load_arg) + # --freeze-all-layers: nothing trains, so load the model in --load weights-only (finetune-style) + # and auto-resume the data position by feeding this run's own progress tracker -- written to + # --save on the previous launch -- into the standard --override-ckpt-iteration path. Reading + # progress from --save (this job's output) rather than --load lets --load stay pinned to a + # fixed checkpoint across resubmits. An explicit --override-ckpt-iteration wins. (The main use + # today is offline-KD teacher-logit dumps.) + if getattr(args, 'freeze_all_layers', False): + # Weights only: don't adopt the loaded checkpoint's optimizer / LR-scheduler / rng, or run + # check_checkpoint_args against a checkpoint from a different run (finetune gates all of + # those; --freeze-all-layers alone would still load the scheduler and assert on arg drift). + args.finetune = True + if args.override_ckpt_iteration is None: + progress_iteration = read_frozen_resume_iteration(args.save) + if progress_iteration > 0: + args.override_ckpt_iteration = progress_iteration + # Finetuning directories pretrained_dir = getattr(args, 'pretrained_checkpoint', None) if pretrained_dir is not None and not checkpoint_exists(load_dir): @@ -2012,43 +2213,58 @@ def load_checkpoint( ) load_dir = pretrained_dir if not checkpoint_exists(load_dir): - raise FileNotFoundError("No checkpoint found in load directory or pretrained directory") + raise FileNotFoundError('No checkpoint found in load directory or pretrained directory') args.finetune = True model = unwrap_model(ddp_model) ckpt_format = args.ckpt_format - if args.auto_detect_ckpt_format or ckpt_format == "torch_dist": + state_dict = None + release = False + if args.auto_detect_ckpt_format or ckpt_format in ('torch_dist', 'fsdp_dtensor'): state_dict, checkpoint_name, release, ckpt_type = _load_base_checkpoint( load_dir, args, rank0=True, checkpointing_context=checkpointing_context ) ckpt_format = None if ckpt_type == CheckpointType.TORCH_DCP: - ckpt_format = "torch_dcp" + ckpt_format = 'torch_dcp' elif ckpt_type == CheckpointType.FSDP_DTENSOR: - ckpt_format = "fsdp_dtensor" + ckpt_format = 'fsdp_dtensor' elif ckpt_type == CheckpointType.LEGACY: - ckpt_format = "torch" + ckpt_format = 'torch' elif ckpt_type in [CheckpointType.LOCAL, CheckpointType.GLOBAL]: - ckpt_format = "torch_dist" + ckpt_format = 'torch_dist' elif ckpt_type == None: pass # Not loaded. else: - raise NotImplementedError(f"checkpoint format {ckpt_format} not supported") + raise NotImplementedError(f'checkpoint format {ckpt_format} not supported') load_kwargs = {} ignore_rng_state = False ignore_rerun_state = True - if ckpt_format == "torch_dist": - ckpt_args = types.SimpleNamespace() - if state_dict is not None and "args" in state_dict: - ckpt_args = state_dict.get("args") + ckpt_args = types.SimpleNamespace() + if ( + ckpt_format in ('torch_dist', 'fsdp_dtensor') + and state_dict is not None + and 'args' in state_dict + ): + ckpt_args = state_dict.get('args') or types.SimpleNamespace() + + # Both model-space torch_dist and fsdp_dtensor checkpoints carry model-keyed + # optimizer state that can be retargeted from GPTModel to HybridModel. + gpt_compat_layer_maps, gpt_compat_load_optim = ( + _maybe_setup_gpt_to_hybrid_load(args, ckpt_args, model) + if ckpt_format in ('torch_dist', 'fsdp_dtensor') and state_dict is not None + else (None, False) + ) + gpt_compat_load_optim = gpt_compat_load_optim and not release - if not hasattr(ckpt_args, "tensor_model_parallel_size"): - print_rank_0("WARNING: TP size not found in checkpoint args, using 1 as default.") - if not hasattr(ckpt_args, "pipeline_model_parallel_size"): - print_rank_0("WARNING: PP size not found in checkpoint args, using 1 as default.") + if ckpt_format == 'torch_dist': + if not hasattr(ckpt_args, 'tensor_model_parallel_size'): + print_rank_0('WARNING: TP size not found in checkpoint args, using 1 as default.') + if not hasattr(ckpt_args, 'pipeline_model_parallel_size'): + print_rank_0('WARNING: PP size not found in checkpoint args, using 1 as default.') ckpt_tp_pp = ( getattr(ckpt_args, "tensor_model_parallel_size", 1), @@ -2060,7 +2276,7 @@ def load_checkpoint( run_world_size = getattr(args, 'world_size', 0) ckpt_dp = getattr(ckpt_args, 'data_parallel_size', 0) run_dp = getattr(args, 'data_parallel_size', 0) - mismatch_msg = "(TP, PP) mismatch after resume ({} vs {} from checkpoint)".format( + mismatch_msg = '(TP, PP) mismatch after resume ({} vs {} from checkpoint)'.format( run_tp_pp, ckpt_tp_pp ) @@ -2087,7 +2303,7 @@ def load_checkpoint( ignore_rng_state = True gen_sd_rng_state = None if ckpt_tp_pp != run_tp_pp: - print_rank_0("{}: RNG state will be ignored".format(mismatch_msg)) + print_rank_0('{}: RNG state will be ignored'.format(mismatch_msg)) if ckpt_type == CheckpointType.LOCAL: sharded_sd_metadata = _build_sharded_state_dict_metadata(args, dp_cp_group=dp_cp_group) @@ -2099,10 +2315,12 @@ def load_checkpoint( f'sharded_state_dict metadata loaded from the checkpoint: {sharded_sd_metadata}' ) - # Determine if optimizer state will be loaded + # Determine if optimizer state will be loaded. For a GPT->hybrid load the + # optimizer state is retargeted at the GPT checkpoint even under --finetune, + # which independently controls iteration and LR-schedule reset semantics. if ( not release - and not args.finetune + and (not args.finetune or gpt_compat_load_optim) and not args.no_load_optim and not getattr(ckpt_args, 'no_save_optim', False) ): @@ -2121,6 +2339,23 @@ def load_checkpoint( else 'dp_zero_gather_scatter' ) } + # Retargeting optimizer state onto a GPT checkpoint only works for the + # model-space formats, where each optimizer ShardedTensor carries the + # model param's key and sharding. Bucket-space formats key state by a + # flat buffer layout that the extra hybrid layers reshuffle. + if gpt_compat_load_optim and sharded_sd_metadata[ + 'distrib_optim_sharding_type' + ] not in ('fully_reshardable', 'fully_sharded_model_space'): + raise RuntimeError( + 'Loading optimizer state from a GPT checkpoint into a hybrid model ' + 'is only supported for model-space distributed-optimizer checkpoints ' + "(sharding type 'fully_reshardable' or 'fully_sharded_model_space'), " + f'but the checkpoint uses ' + f'{sharded_sd_metadata["distrib_optim_sharding_type"]!r}. Re-save the ' + 'GPT checkpoint with --dist-ckpt-optim-fully-reshardable, or pass ' + '--no-load-optim ' + 'to start from a fresh optimizer.' + ) if ( ckpt_tp_pp != run_tp_pp and sharded_sd_metadata['distrib_optim_sharding_type'] @@ -2153,7 +2388,7 @@ def load_checkpoint( # Ensure we have a dict before updating to avoid NoneType AttributeError. if sharded_sd_metadata is None: sharded_sd_metadata = {} - sharded_sd_metadata["dp_cp_group"] = dp_cp_group + sharded_sd_metadata['dp_cp_group'] = dp_cp_group optim_sd_kwargs = dict(metadata=sharded_sd_metadata, is_loading=True) model_sd_kwargs = dict(metadata=sharded_sd_metadata) @@ -2196,26 +2431,42 @@ def load_checkpoint( model_sd_kwargs=model_sd_kwargs, rerun_state=gen_sd_rerun_state, ) - elif args.ckpt_format == "torch_dcp": + + if gpt_compat_layer_maps is not None: + from megatron.core.dist_checkpointing.gpt_checkpoint_interop import ( + retarget_sharded_state_dict_to_gpt_checkpoint, + ) + + # The optimizer sharded state dict is built from the (hybrid) model sharded + # state dict, so its entries carry the same ``decoder.layers..`` keys and + # sharding; the same retargeting points them at the GPT checkpoint too. + for sd_key, sub_sd in load_kwargs['sharded_state_dict'].items(): + is_model = sd_key == 'model' or re.fullmatch(r'model\d+', sd_key) + is_optim = gpt_compat_load_optim and ( + sd_key == 'optimizer' or re.fullmatch(r'optimizer\d+', sd_key) + ) + if is_model or is_optim: + retarget_sharded_state_dict_to_gpt_checkpoint(sub_sd, gpt_compat_layer_maps) + elif args.ckpt_format == 'torch_dcp': model_sd = model[0].state_dict() optimizer_sd = optimizer.state_dict(is_loading=True) if tp_group is None and pp_group is None: tp_group = mpu.get_tensor_model_parallel_group() pp_group = mpu.get_pipeline_model_parallel_group() sharded_state_dict = { - "model": model_sd, - "optimizer": optimizer_sd, - "args": None, - "iteration": 1, - "rng_state": get_rng_state( + 'model': model_sd, + 'optimizer': optimizer_sd, + 'args': None, + 'iteration': 1, + 'rng_state': get_rng_state( args.ckpt_format, tp_group, pp_group, dp_cp_group=dp_cp_group, dp_group=dp_group ), - "checkpoint_version": None, - "opt_param_scheduler": opt_param_scheduler.state_dict(), - "num_floating_point_operations_so_far": 0, + 'checkpoint_version': None, + 'opt_param_scheduler': opt_param_scheduler.state_dict(), + 'num_floating_point_operations_so_far': 0, } - load_kwargs["sharded_state_dict"] = sharded_state_dict - elif args.ckpt_format == "fsdp_dtensor": + load_kwargs['sharded_state_dict'] = sharded_state_dict + elif args.ckpt_format == 'fsdp_dtensor': reader = FileSystemReader(get_load_checkpoint_path_by_args(args)) try: state_dict_metadata = reader.read_metadata().state_dict_metadata @@ -2227,7 +2478,7 @@ def load_checkpoint( gen_sd_rng_state = None gen_sd_optim = None if not args.finetune: - if "rerun_state_machine" in state_dict_metadata: + if 'rerun_state_machine' in state_dict_metadata: gen_sd_rerun_state = get_rerun_state_machine().state_dict( data_iterator=None, ckpt_format=ckpt_format, force=True ) @@ -2235,8 +2486,9 @@ def load_checkpoint( gen_sd_rng_state = get_rng_state( args.ckpt_format, tp_group, pp_group, dp_cp_group=dp_cp_group, dp_group=dp_group ) - if not args.no_load_optim: - gen_sd_optim = optimizer + if (not args.finetune or gpt_compat_load_optim) and not args.no_load_optim: + gen_sd_optim = optimizer + if not args.finetune: gen_sd_opt_param_scheduler = opt_param_scheduler optim_sd_kwargs = dict( @@ -2244,18 +2496,43 @@ def load_checkpoint( is_loading=True, ) - state_dict = generate_state_dict( - args, - model=model, - optimizer=gen_sd_optim, - opt_param_scheduler=gen_sd_opt_param_scheduler, - rng_state=gen_sd_rng_state, - optim_sd_kwargs=optim_sd_kwargs, - rerun_state=gen_sd_rerun_state, - iteration=1, - ) - state_dict["_model"] = model - load_kwargs["sharded_state_dict"] = state_dict + # Megatron-FSDP materializes optimizer slots with a dummy zero-gradient + # step while building a loading state dict. A normal full resume + # overwrites every model parameter afterward, but GPT->Hybrid leaves + # fresh-only layers untouched. Temporarily zero the optimizer LR so the + # shape-materialization step cannot weight-decay or momentum-update + # those fresh model parameters. + optimizer_lrs = [] + if gpt_compat_layer_maps is not None and gen_sd_optim is not None: + pending_optimizers = [gen_sd_optim] + while pending_optimizers: + pending_optimizer = pending_optimizers.pop() + if hasattr(pending_optimizer, 'chained_optimizers'): + pending_optimizers.extend(pending_optimizer.chained_optimizers) + continue + inner_optimizer = getattr(pending_optimizer, 'optimizer', None) + if inner_optimizer is None: + continue + for param_group in inner_optimizer.param_groups: + optimizer_lrs.append((param_group, param_group.get('lr'))) + param_group['lr'] = 0.0 + + try: + state_dict = generate_state_dict( + args, + model=model, + optimizer=gen_sd_optim, + opt_param_scheduler=gen_sd_opt_param_scheduler, + rng_state=gen_sd_rng_state, + optim_sd_kwargs=optim_sd_kwargs, + rerun_state=gen_sd_rerun_state, + iteration=1, + ) + finally: + for param_group, lr in optimizer_lrs: + param_group['lr'] = lr + state_dict['_model'] = model + load_kwargs['sharded_state_dict'] = state_dict state_dict, checkpoint_name, release, ckpt_type = _load_base_checkpoint( load_dir, @@ -2264,6 +2541,7 @@ def load_checkpoint( checkpointing_context=checkpointing_context, dp_cp_group=dp_cp_group, expt_dp_group=expt_dp_group, + gpt_compat_layer_maps=gpt_compat_layer_maps, **load_kwargs, ) @@ -2276,9 +2554,9 @@ def load_checkpoint( set_checkpoint_version(state_dict.get('checkpoint_version', 0)) # Convert to regular torch tensor to DTensor. - if ckpt_type == CheckpointType.LEGACY and args.ckpt_format == "torch_dcp": - dtensor_state_dict = _to_dtensor(ddp_model, state_dict["model"]) - state_dict["model"] = dtensor_state_dict + if ckpt_type == CheckpointType.LEGACY and args.ckpt_format == 'torch_dcp': + dtensor_state_dict = _to_dtensor(ddp_model, state_dict['model']) + state_dict['model'] = dtensor_state_dict # Set iteration. if args.finetune or release: @@ -2300,7 +2578,12 @@ def load_checkpoint( # Check arguments. if 'args' in state_dict and not args.finetune: checkpoint_args = state_dict['args'] - check_checkpoint_args(checkpoint_args) + # A GPT block is split into separate attention and MLP positions in + # HybridModel, so num_layers intentionally differs even for an + # architecture-preserving load. Keep every other resume-time argument + # compatibility check. + skip_args = {'num_layers'} if gpt_compat_layer_maps is not None else None + check_checkpoint_args(checkpoint_args, skip_args=skip_args) args.consumed_train_samples = getattr(checkpoint_args, 'consumed_train_samples', 0) args.skipped_train_samples = getattr(checkpoint_args, 'skipped_train_samples', 0) update_num_microbatches(consumed_samples=args.consumed_train_samples, verbose=True) @@ -2308,18 +2591,66 @@ def load_checkpoint( else: print_rank_0('could not find arguments in the checkpoint ...') + # --override-ckpt-iteration: rewind the data loader to this iteration, operating on `args` + # (not state_dict) so it also works on checkpoints with no saved `args` (release / HF). The + # GBS-match check applies only when adopting this checkpoint's own args (not a finetune load). + if getattr(args, 'override_ckpt_iteration', None) is not None: + if 'args' in state_dict and not args.finetune: + ckpt_global_batch_size = getattr(state_dict['args'], 'global_batch_size', None) + if ( + ckpt_global_batch_size is not None + and ckpt_global_batch_size != args.global_batch_size + ): + warn_rank_0( + '--override-ckpt-iteration recomputes consumed_train_samples = target_iter * ' + f'current global_batch_size ({args.global_batch_size}), but this checkpoint ' + f'was saved at global_batch_size {ckpt_global_batch_size}. If the target ' + "iteration is in the checkpoint run's units, the data loader resumes from a " + 'shifted sample offset.' + ) + iteration = args.override_ckpt_iteration + args.consumed_train_samples = iteration * args.global_batch_size + args.skipped_train_samples = 0 + update_num_microbatches(consumed_samples=args.consumed_train_samples, verbose=True) + print_rank_0( + f'--override-ckpt-iteration: start at iteration {iteration} ' + f'(consumed_train_samples {args.consumed_train_samples})' + ) + def load_model_state_dict(module, state_dict, strict: bool): """Helper function to load state dict with fallback for missing extra states.""" + # GTP native-FP8 weights: load_state_dict's copy_ re-quantizes into the FP8 param, which + # TE's IsMXFP8Tensor check rejects for our subclass. Present the base FP8 class for it. + from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + + if HAVE_GTP: + from megatron.core.tensor_parallel.gtp_api import gtp_native_fp8_load_context + + load_ctx = lambda: gtp_native_fp8_load_context(module) + else: + from contextlib import nullcontext + + load_ctx = nullcontext try: - module.load_state_dict(state_dict, strict=strict) + with load_ctx(): + module.load_state_dict(state_dict, strict=strict) except Exception as e: if strict: # Fallback support for backward compatibility breaking changes in TransformerEngine - load_return = module.load_state_dict(state_dict, strict=False) - print(f"load_return: {load_return}") + with load_ctx(): + load_return = module.load_state_dict(state_dict, strict=False) + print(f'load_return: {load_return}') + + # Megatron-FSDP DTensors are loaded into the model buffers in-place above. + # Replaying the translated raw state dict through ``load_state_dict`` would + # reapply HybridModel-local entries that were intentionally omitted (for + # example fresh Mamba layers). + gpt_fsdp_model_loaded_in_place = ( + ckpt_format == 'fsdp_dtensor' and gpt_compat_layer_maps is not None + ) # Model. - if not skip_load_to_model_and_opt: + if not skip_load_to_model_and_opt and not gpt_fsdp_model_loaded_in_place: if len(ddp_model) == 1: load_model_state_dict(ddp_model[0], state_dict['model'], strict) else: @@ -2335,7 +2666,7 @@ def load_model_state_dict(module, state_dict, strict: bool): fix_query_key_value_ordering(model, checkpoint_version) # Optimizer. - if not release and not args.finetune and not args.no_load_optim: + if not release and (not args.finetune or gpt_compat_load_optim) and not args.no_load_optim: try: # Load state dict. if ( @@ -2377,8 +2708,8 @@ def load_model_state_dict(module, state_dict, strict: bool): update_legacy_format=args.ckpt_convert_update_legacy_dist_opt_format, ) - # Load scheduler. - if opt_param_scheduler is not None: + # Load scheduler unless --finetune requests a fresh iteration and LR schedule. + if opt_param_scheduler is not None and not args.finetune: if 'lr_scheduler' in state_dict: # backward compatbility opt_param_scheduler.load_state_dict(state_dict['lr_scheduler']) else: @@ -2404,7 +2735,7 @@ def load_model_state_dict(module, state_dict, strict: bool): if 'rerun_state_machine' in state_dict: get_rerun_state_machine().load_state_dict(state_dict['rerun_state_machine']) except Exception as e: - print_rank_0(f"Unable to restore RerunMachine from checkpoint: {e}. Skipping.") + print_rank_0(f'Unable to restore RerunMachine from checkpoint: {e}. Skipping.') # rng states. if not release and not args.finetune and not args.no_load_rng and not ignore_rng_state: @@ -2412,7 +2743,7 @@ def load_model_state_dict(module, state_dict, strict: bool): cuda_rng_tracker = tensor_parallel.get_cuda_rng_tracker() graph_safe_rng = tensor_parallel.is_graph_safe_cuda_rng_tracker(cuda_rng_tracker) if 'rng_state' in state_dict: - if args.ckpt_format == "fsdp_dtensor": + if args.ckpt_format == 'fsdp_dtensor': # FSDP DTensor checkpoints store rng_state in a different format. tp_rank = ( get_pg_rank(tp_group) @@ -2427,7 +2758,7 @@ def load_model_state_dict(module, state_dict, strict: bool): if f"({pp_rank}, {tp_rank})" in state_dict['rng_state']: rng_state = state_dict['rng_state'][f"({pp_rank}, {tp_rank})"] else: - print_rank_0("WARNING: RNG state not found for current TP/PP rank") + print_rank_0('WARNING: RNG state not found for current TP/PP rank') rng_state = next(iter(state_dict['rng_state'].values())) else: rng_state = state_dict['rng_state'] @@ -2493,9 +2824,12 @@ def load_model_state_dict(module, state_dict, strict: bool): if pp_group is not None else mpu.get_pipeline_model_parallel_world_size() ) + _gtp_remat_r = mpu.get_gtp_weight_remat_rank() + _gtp_remat_w = mpu.get_gtp_weight_remat_world_size() print_rank_0( f' successfully loaded checkpoint from {load_dir} ' f'[ t {_tp_r + 1}/{_tp_w}, ' + f'gtp_remat {_gtp_remat_r + 1}/{_gtp_remat_w}, ' f'p {_pp_r + 1}/{_pp_w} ] ' f'at iteration {iteration}' ) @@ -2524,28 +2858,7 @@ def load_model_state_dict(module, state_dict, strict: bool): log_printed = True if has_nvidia_modelopt: - print_distributed_quant_summary(model, msg="After loading checkpoint") - - # Load teacher model in Distillation mode. - if getattr(args, "export_kd_teacher_load", None): - from megatron.post_training.checkpointing import load_modelopt_checkpoint - - unwrapped_model = unwrap_model(model)[0] - # Note: load_modelopt_checkpoint may call this function so we prevent infinite recursion. - if hasattr(unwrapped_model, 'teacher_model'): - teacher = unwrapped_model.teacher_model - print_rank_0( - f"Loading teacher as {type(teacher).__name__} from {args.export_kd_teacher_load} ..." - ) - # [WAR]: To avoid error out on loading teacher's checkpoint, we temporarily - # set args.finetune to True while loading the teacher checkpoint. - original_args_finetune, original_ckpt_format = args.finetune, args.ckpt_format - args.finetune = True - if args.export_kd_teacher_ckpt_format is not None: - args.ckpt_format = args.export_kd_teacher_ckpt_format - load_modelopt_checkpoint([teacher], load_arg='export_kd_teacher_load') - args.finetune, args.ckpt_format = original_args_finetune, original_ckpt_format - print_rank_0("... teacher loaded successfully.") + print_distributed_quant_summary(model, msg='After loading checkpoint') return iteration, num_floating_point_operations_so_far @@ -2556,7 +2869,7 @@ def _to_dtensor(wrapped_model, model_state_dict): new_model_sd = dict() for k, v in model_state_dict.items(): # FP8 extra state cannot be converted to dtensor yet. - if "_extra_state" in k: + if '_extra_state' in k: new_model_sd[k] = v else: new_model_sd[k] = torch.distributed.tensor.distribute_tensor(v, device_mesh) @@ -2580,7 +2893,7 @@ def load_biencoder_checkpoint( tracker_filename = get_checkpoint_tracker_filename(load_path) - with open_file(tracker_filename, 'r') as f: + with maybe_msc.open(tracker_filename, 'r') as f: iteration = int(f.read().strip()) checkpoint_name = get_checkpoint_name( diff --git a/megatron/training/config/container.py b/megatron/training/config/container.py index d872701290b..39528c758be 100644 --- a/megatron/training/config/container.py +++ b/megatron/training/config/container.py @@ -14,7 +14,7 @@ HAVE_YAML = False from megatron.core.distributed.distributed_data_parallel_config import DistributedDataParallelConfig -from megatron.core.msc_utils import MultiStorageClientFeature +from megatron.core.msc_utils import maybe_msc from megatron.core.optimizer import OptimizerConfig from megatron.training.config.common_config import DistributedInitConfig, ProfilingConfig, RNGConfig from megatron.training.config.inference_config import InferenceSetupConfig @@ -113,22 +113,13 @@ def from_yaml( from omegaconf import OmegaConf - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - yaml_path_exists = msc.os.path.exists(yaml_path) - else: - yaml_path_exists = os.path.exists(yaml_path) + yaml_path_exists = maybe_msc.os.path.exists(yaml_path) if not yaml_path_exists: raise FileNotFoundError(f"YAML file not found: {yaml_path}") - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - with msc.open(yaml_path, "r") as f: - config_dict = yaml.safe_load(f) - else: - with open(yaml_path, "r") as f: - config_dict = yaml.safe_load(f) + with maybe_msc.open(yaml_path, "r") as f: + config_dict = yaml.safe_load(f) # Convert to OmegaConf first for better compatibility with instantiate conf = OmegaConf.create(config_dict) @@ -222,13 +213,8 @@ def to_yaml(self, yaml_path: str) -> None: config_dict = self.to_dict() with safe_yaml_representers(): - if MultiStorageClientFeature.is_enabled(): - msc = MultiStorageClientFeature.import_package() - with msc.open(yaml_path, "w") as f: - yaml.safe_dump(config_dict, f, default_flow_style=False) - else: - with open(yaml_path, "w") as f: - yaml.safe_dump(config_dict, f, default_flow_style=False) + with maybe_msc.open(yaml_path, "w") as f: + yaml.safe_dump(config_dict, f, default_flow_style=False) def print_yaml(self) -> None: """ diff --git a/megatron/training/config/inference_config.py b/megatron/training/config/inference_config.py index edad1f4d21d..017ba3966c1 100644 --- a/megatron/training/config/inference_config.py +++ b/megatron/training/config/inference_config.py @@ -132,14 +132,19 @@ class InferenceSetupConfig: down to tp_size, giving a log-spaced distribution with bounded relative padding. "linear" uses varying linear strides across the range.""" - inference_dynamic_batching_sampling_backend: Literal["torch", "flashinfer"] = "torch" - """Which sampling kernels to use during inference. Falls back to "torch" with a warning if - "flashinfer" is requested but the package is not installed.""" + inference_dynamic_batching_sampling_backend: Literal["torch", "flashinfer"] = "flashinfer" + """Which sampling kernels to use during inference. Defaults to "flashinfer" and falls back to + "torch" with a warning if the flashinfer package is not installed.""" - inference_dynamic_batching_async_sched_mode: Literal["legacy", "serial"] = "legacy" + offset_sampling_seed_by_dp_rank: bool = True + """Offset the inference sampling seed by the data-parallel rank so each DP rank gets a unique + generation seed. Disable with --use-same-sampling-seed-across-dp-ranks. Also forced off when + --deterministic-mode is enabled.""" + + inference_dynamic_batching_async_sched_mode: Literal["legacy", "async"] = "legacy" """Async scheduling mode for dynamic batching. "legacy" (default) preserves the - existing resolve-before-prepare path. "serial" speculatively prepares and forwards decode-only - steps before resolving finished requests.""" + existing resolve-before-prepare path. "async" overlaps asynchronous scheduling phases by + reordering them to prepare-before-resolve.""" inference_dynamic_batching_logprobs_mode: Literal["raw_logprobs", "processed_logprobs"] = ( "raw_logprobs" @@ -156,6 +161,9 @@ class InferenceSetupConfig: """Extend prefill/mixed CUDA graph capture up to `max_tokens`. By default, all graphs are limited by the decode limit of `max_requests * (num_speculative_tokens + 1)`.""" + inference_cuda_graph_max_tokens: int = 512 + """Token ceiling for the largest captured prefill/mixed CUDA graph (default: 512).""" + # ---------------- Chunked prefill / speculation ---------------- enable_chunked_prefill: bool = False @@ -176,11 +184,14 @@ class InferenceSetupConfig: space is needed.""" inference_dynamic_batching_prefix_caching_coordinator_policy: Literal[ - "longest_prefix", "first_prefix_block", "round_robin" - ] = "first_prefix_block" - """Coordinator routing policy for prefix caching. "first_prefix_block" (default) routes based on - the first block hash only. "longest_prefix" routes to the rank with the longest matching prefix. - "round_robin" ignores prefix affinity and cycles through ranks.""" + "longest_prefix", "first_prefix_block", "load_balanced" + ] = "load_balanced" + """Coordinator routing policy for prefix caching. "load_balanced" (default) routes to the rank + with the fewest in-flight requests, ignoring prefix affinity. "first_prefix_block" routes based + on the first block hash only. "longest_prefix" routes to the rank with the longest matching + prefix. "first_prefix_block" and "longest_prefix" both combine prefix affinity with load + balancing and fall back to load-balanced routing when prefix caching is disabled or no prefix + match exists.""" inference_dynamic_batching_prefix_caching_routing_alpha: float = 0.5 """Weight for prefix-aware routing score: score = alpha * match + (1 - alpha) * normalized_load. @@ -337,6 +348,7 @@ def to_inference_config( ), use_cuda_graphs_for_non_decode_steps=not self.decode_only_cuda_graphs, cuda_graph_all_prefills=self.inference_cuda_graph_all_prefills, + cuda_graph_max_tokens=self.inference_cuda_graph_max_tokens, static_kv_memory_pointers=static_kv_memory_pointers, max_sequence_length=max_sequence_length, mamba_inference_state_config=mamba_inference_state_config, @@ -365,6 +377,7 @@ def to_inference_config( use_synchronous_zmq_collectives=self.inference_use_synchronous_zmq_collectives, disable_ep_consensus=self.inference_disable_ep_consensus, sampling_backend=self.inference_dynamic_batching_sampling_backend, + offset_sampling_seed_by_dp_rank=self.offset_sampling_seed_by_dp_rank, async_sched_mode=AsyncScheduleMode( self.inference_dynamic_batching_async_sched_mode ), diff --git a/megatron/training/config/training_config.py b/megatron/training/config/training_config.py index 294e28f6dfd..6df0f4bb21c 100644 --- a/megatron/training/config/training_config.py +++ b/megatron/training/config/training_config.py @@ -660,6 +660,13 @@ class CheckpointConfig: verify_integrity: bool = False """Whether to hash checkpointing files during save and validate their integrity during load.""" + stream_ckpt_dequant: bool = True + """Per-tensor streaming dequantize when loading checkpoints with quantized model params + (FP8, MXFP8, blockwise FP8, NVFP4). The LoadPlanner dequantizes one destination at a time, + instead of dequantizing the entire state dict to high precision before the load starts + (which allocates N simultaneous scratch tensors and can OOM on large models). On by + default; pass --no-stream-ckpt-dequant to fall back to the legacy upfront pass.""" + def __post_init__(self): from megatron.training.utils import has_nvrx_checkpointing_async_support diff --git a/megatron/training/datasets/data_samplers.py b/megatron/training/datasets/data_samplers.py index 27ffefa19c3..fd0df9f8b8b 100644 --- a/megatron/training/datasets/data_samplers.py +++ b/megatron/training/datasets/data_samplers.py @@ -20,6 +20,12 @@ def build_pretraining_data_loader(dataset, consumed_samples): if dataset is None: return None + # Empty split (e.g. valid/test when --eval-iters 0): return null loader + try: + if len(dataset) == 0: + return None + except TypeError: + pass args = get_args() if hasattr(dataset, 'split'): diff --git a/megatron/training/distillation/utils_logits.py b/megatron/training/distillation/utils_logits.py index 07b75f44a49..d4602ce27ab 100644 --- a/megatron/training/distillation/utils_logits.py +++ b/megatron/training/distillation/utils_logits.py @@ -32,7 +32,7 @@ except ImportError: HAVE_ZSTANDARD = False -from megatron.core.msc_utils import MultiStorageClientFeature +from megatron.core.msc_utils import MultiStorageClientFeature, maybe_msc from megatron.training import get_args from megatron.training.utils import get_blend_and_blend_per_split @@ -94,11 +94,7 @@ def storage_makedirs(path: str, exist_ok: bool = True) -> None: """Create a local or MSC directory/prefix.""" if not path: return - msc = _msc_if_needed(path) - if msc is not None: - msc.os.makedirs(path, exist_ok=exist_ok) - else: - os.makedirs(path, exist_ok=exist_ok) + maybe_msc.os.makedirs(path, exist_ok=exist_ok) def storage_move(src: str, dst: str) -> None: diff --git a/megatron/training/global_vars.py b/megatron/training/global_vars.py index e08667bdd91..ba261fad8f1 100644 --- a/megatron/training/global_vars.py +++ b/megatron/training/global_vars.py @@ -136,7 +136,8 @@ def set_global_variables(args, build_tokenizer=True): rank=args.rank, global_batch_size=args.global_batch_size, micro_batch_size=args.micro_batch_size, - data_parallel_size=args.data_parallel_size, + # Full DP x gtp_remat degree (args.data_parallel_size is the gtp_remat-excluded replicate). + data_parallel_size=args.data_parallel_size * args.gtp_weight_remat_size, decrease_batch_size_if_needed=args.decrease_batch_size_if_needed, step_batch_size_schedule=args.step_batch_size_schedule, seq_length=args.seq_length, diff --git a/megatron/training/initialize.py b/megatron/training/initialize.py index 40c5b8ad11f..4d49cebfae1 100644 --- a/megatron/training/initialize.py +++ b/megatron/training/initialize.py @@ -360,12 +360,24 @@ def _initialize_distributed(get_embedding_ranks, get_position_embedding_ranks, s if mpu.model_parallel_is_initialized(): print("model parallel is already initialized") else: + if args.gtp_weight_remat_size > 1 or args.expert_gtp_weight_remat_size > 1: + from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + + assert HAVE_GTP, ( + "GTP requires TransformerEngine >= 2.19. " + "Set both --gtp_remat-weight-remat-size and " + "--expert-generalized-tensor-parallel-remat-size to 1 to disable GTP." + ) mpu.initialize_model_parallel( args.tensor_model_parallel_size, args.pipeline_model_parallel_size, args.virtual_pipeline_model_parallel_size, pipeline_model_parallel_comm_backend=args.pipeline_model_parallel_comm_backend, use_sharp=args.use_sharp, + # GTP_remat/EGTP_remat need world divisible by TP*PP*CP*GTP_remat (expert grid + # by ETP*EP*PP*EGTP_remat). Inactive when the remat sizes are 1. + gtp_remat_size=args.gtp_weight_remat_size, + expert_gtp_remat_size=args.expert_gtp_weight_remat_size, context_parallel_size=args.context_parallel_size, hierarchical_context_parallel_sizes=args.hierarchical_context_parallel_sizes, dynamic_context_parallel=args.dynamic_context_parallel, @@ -386,6 +398,10 @@ def _initialize_distributed(get_embedding_ranks, get_position_embedding_ranks, s f"> initialized tensor model parallel with size " f"{mpu.get_tensor_model_parallel_world_size()}" ) + print_rank_0( + f"> initialized gtp weight remat with size " + f"{mpu.get_gtp_weight_remat_world_size()}" + ) print_rank_0( f"> initialized pipeline model parallel with size " f"{mpu.get_pipeline_model_parallel_world_size()}" diff --git a/megatron/training/models/dist_utils.py b/megatron/training/models/dist_utils.py index 3d4d7e63f96..933d615e422 100644 --- a/megatron/training/models/dist_utils.py +++ b/megatron/training/models/dist_utils.py @@ -1,9 +1,6 @@ # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. import logging - -logger = logging.getLogger(__name__) - from typing import Any, Callable import torch @@ -14,6 +11,7 @@ DistributedDataParallelConfig, FullyShardedDataParallel, ) +from megatron.core.full_cuda_graph import get_shared_capture_stream from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer from megatron.core.optimizer.layer_wise_optimizer import ( LayerWiseDistributedOptimizer, @@ -31,7 +29,7 @@ from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer import MegatronModule, TransformerConfig from megatron.core.transformer.module import Float16Module -from megatron.core.utils import get_model_config +from megatron.core.utils import get_model_config, get_pg_rank try: from megatron.core.fp8_utils import correct_amax_history_if_needed @@ -39,6 +37,9 @@ correct_amax_history_if_needed = None +logger = logging.getLogger(__name__) + + def unimodal_build_distributed_models( build_model_func: Callable, transformer_config: TransformerConfig, @@ -217,7 +218,7 @@ def _print_num_params(model: list[MegatronModule], pg_collection: ProcessGroupCo """Print the number of parameters in the model on rank 0. Only prints on data parallel rank 0 to avoid duplicate output. - Shows parameter count per (tensor parallel, pipeline parallel) rank. + Shows parameter count per (tensor parallel, gtp_remat, pipeline parallel) rank. Args: model: List of model modules to count parameters from @@ -225,8 +226,9 @@ def _print_num_params(model: list[MegatronModule], pg_collection: ProcessGroupCo """ if (pg_collection.dp.rank() == 0) and (pg_collection.cp.rank() == 0): print( - " > number of parameters on (tensor, pipeline) model parallel rank ({}, {}): {}".format( + " > number of parameters on (tensor, gtp_remat, pipeline) model parallel rank ({}, {}, {}): {}".format( pg_collection.tp.rank(), + get_pg_rank(pg_collection.gtp_remat), pg_collection.pp.rank(), sum( [ @@ -331,11 +333,15 @@ def _ddp_wrap( if not ddp_config.overlap_grad_reduce: ddp_config.bucket_size = None - # DDP initialization is required to be on a side-stream for the full-iteration CUDA graph. - # this side-stream may be nested if being called from within the get_model function, but it - # is here in case someone wants to use this directly outside of get_model. - ddp_stream = torch.cuda.Stream() + if get_model_config(model[0]).cuda_graph_impl == "full_iteration": + # DDP initialization must use the full-iteration capture stream so its retained + # AccumulateGrad nodes do not reference a different, non-capturing stream. + ddp_stream = get_shared_capture_stream() + else: + # Preserve a dedicated initialization stream for all other implementations. + ddp_stream = torch.cuda.Stream() ddp_stream.wait_stream(torch.cuda.current_stream()) + with torch.cuda.stream(ddp_stream): dp_init_kwargs = {} if not use_torch_fsdp2: @@ -354,12 +360,24 @@ def _ddp_wrap( effective_bucket_size = ( None if disable_bucketing or pp_rank > 0 else ddp_config.bucket_size ) + # Size the layout by the group the optimizer actually shards over, which is + # the intra-instance group when there are several optimizer instances. Using + # the full dp_cp would report more shards than the reduce-scatter uses and + # leave the trailing shard of every bucket owned by no rank. + intra_dp_cp_group = getattr(pg_collection, "intra_dp_cp", None) + intra_expt_dp_group = getattr(pg_collection, "intra_expt_dp", None) chunk_kwargs["full_param_layout"] = compute_layout( all_params, effective_bucket_size, - pg_collection.dp_cp.size(), + ( + intra_dp_cp_group if intra_dp_cp_group is not None else pg_collection.dp_cp + ).size(), ddp_config, - expert_data_parallel_world_size=pg_collection.expt_dp.size(), + expert_data_parallel_world_size=( + intra_expt_dp_group + if intra_expt_dp_group is not None + else pg_collection.expt_dp + ).size(), ) wrapped_chunk = DP( @@ -372,7 +390,7 @@ def _ddp_wrap( wrapped_model.append(wrapped_chunk) model = wrapped_model - # Critical: ensure side-stream work completes before touching params on default stream + # Ensure initialization-stream work completes before touching params on the default stream. torch.cuda.current_stream().wait_stream(ddp_stream) # Broadcast params from data parallel src rank to other data parallel ranks. diff --git a/megatron/training/training.py b/megatron/training/training.py index 4ad7f00c4db..feb7855122d 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -42,7 +42,7 @@ # First-party. from megatron.core import mpu, nccl_allocator, tensor_parallel -from megatron.core.datasets.data_schedule import wrap_data_iterator +from megatron.core.datasets.data_schedule import HybridCPDataLoaderWrapper from megatron.core.distributed import DistributedDataParallel as DDP from megatron.core.distributed import ( DistributedDataParallelConfig, @@ -54,13 +54,13 @@ ) from megatron.core.enums import ModelType from megatron.core.fp8_utils import correct_amax_history_if_needed -from megatron.core.full_cuda_graph import FullCudaGraphWrapper +from megatron.core.full_cuda_graph import FullCudaGraphWrapper, get_shared_capture_stream from megatron.core.inference.symmetric_memory import SymmetricMemoryManager from megatron.core.inference.unified_memory import create_unified_mempool from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( is_linear_attention_variant, ) -from megatron.core.msc_utils import MultiStorageClientFeature, open_file +from megatron.core.msc_utils import maybe_msc from megatron.core.num_microbatches_calculator import ( destroy_num_microbatches_calculator, get_current_global_batch_size, @@ -114,6 +114,7 @@ get_rerun_state_machine, ) from megatron.core.resharding.refit import swap_model_weights +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP from megatron.core.transformer.cuda_graphs import TECudaGraphHelper from megatron.core.transformer.experimental_attention_variant.dsa import DSAIndexerLossLoggingHelper from megatron.core.transformer.module import Float16Module @@ -149,7 +150,7 @@ set_jit_fusion_options, write_args_to_tensorboard, ) -from megatron.training.utils import is_hybrid_model +from megatron.training.utils import is_gtp_remat_active, is_hybrid_model # Local. from . import ft_integration, one_logger_utils @@ -207,6 +208,8 @@ try: from modelopt.torch.distill.plugins.megatron import get_tensor_shapes_adjust_fn_for_distillation + from megatron.post_training.utils import maybe_enable_modelopt + has_nvidia_modelopt = True except ImportError: has_nvidia_modelopt = False @@ -384,116 +387,6 @@ def consume_seqlen_stats_in_iteration() -> Tuple[Optional[float], Optional[float return total_real_tokens / dedup, seqlen_squared_sum / dedup -def _dsv4_hybrid_self_attention_flops( - *, - hidden_size, - num_attention_heads, - v_head_dim, - q_lora_rank, - o_groups, - o_lora_rank, - csa_window_size, - seq_length, - n_layers_r0, - n_layers_r4, - n_layers_r128, - dsa_indexer_n_heads, - dsa_indexer_head_dim, - dsa_indexer_topk, -): - """Per-iteration DSv4-hybrid self-attention FLOPs coefficients. - - DSv4 attention layers are MLA-based but replace full ``O(L^2)`` core - attention with sparse attention (sliding window + compressed-KV), plus a - main compressor and, on ``ratio==4`` layers, a learned indexer (DSA). This - is the SINGLE SOURCE OF TRUTH shared by the standard-model path - (``transformer_flops``) and the hybrid-model path (``hybrid_flops``): the - two differ only in how they obtain the per-ratio attention-layer counts - (a ``csa_compress_ratios`` list vs. Window/CSA/HCA pattern symbols). - - Layer-type ratios: - * ``r==0`` (Window): window-only attention; no compressor / indexer. - * ``r==4`` (CSA): window + learned-topk over compressed KV - (compressor + indexer). - * ``r==128`` (HCA): window + all compressed KV (compressor only). - - Returns ``(token_linear, core)`` WITHOUT the fwd+bwd (x3) or FMA (x2) - expansion factors -- the caller applies those. Multiply ``token_linear`` by - the real (unpadded) token count and ``core`` by ``sum_i(L_i ** 2)``. - """ - n_attn_layers = n_layers_r0 + n_layers_r4 + n_layers_r128 - - # ---- MLA projections (token-linear, per attention layer) ---- - # DSv4 hybrid MLA: ``qk_head_dim + qk_pos_emb_head_dim == v_head_dim`` and the - # joint KV is a single ``hidden -> v_head_dim`` projection; the output uses a - # grouped low-rank ``(o_groups x o_lora_rank)`` projection (NOT the dense - # ``num_heads * v_head_dim -> hidden`` projection of plain MLA). - q_term = q_lora_rank * (hidden_size + num_attention_heads * v_head_dim + 1) - kv_term = hidden_size * v_head_dim + v_head_dim - o_term = num_attention_heads * v_head_dim * o_lora_rank + o_groups * o_lora_rank * hidden_size - mla_proj_term = (q_term + kv_term + o_term) * n_attn_layers - - # ---- Sparse attention (replaces full core attention) ---- - # Split into token-linear parts (window attention, constant per token) and - # L^2 parts (compressed-KV attention, scales with sequence length) so THD - # packed sequences get correct ``seqlen_squared_sum`` scaling. - # r=0: window-only, fixed per-token cost. - sparse_attn_r0 = n_layers_r0 * num_attention_heads * csa_window_size * v_head_dim * 2 - # r=128: window (token-linear) + all compressed KV (L^2). Compressed - # positions per token ~= L/(128*2) (causal /2), x2 for QK^T + softmax@V -> - # L^2 coefficient ``num_heads * v_head_dim / 128``. - sparse_attn_r128_window = n_layers_r128 * num_attention_heads * csa_window_size * v_head_dim * 2 - sparse_attn_r128_core = n_layers_r128 * num_attention_heads * v_head_dim / 128 - - # ---- Main compressor (ratio > 0 layers) ---- - # Two projections per layer (wkv + wgate): ``hidden -> coff * v_head_dim``. - # ratio == 4: coff = 2 (overlapping windows); ratio == 128: coff = 1. - main_compressor_term = ( - n_layers_r4 * hidden_size * (2 * v_head_dim) * 2 - + n_layers_r128 * hidden_size * (1 * v_head_dim) * 2 - ) - - # ---- r=4 layers: sparse attention + indexer ---- - if n_layers_r4 > 0: - assert ( - dsa_indexer_n_heads is not None - ), "dsa_indexer_n_heads must be set for dsv4_hybrid with ratio==4 layers." - assert ( - dsa_indexer_head_dim is not None - ), "dsa_indexer_head_dim must be set for dsv4_hybrid with ratio==4 layers." - assert ( - dsa_indexer_topk is not None - ), "dsa_indexer_topk must be set for dsv4_hybrid with ratio==4 layers." - effective_topk_4 = min(dsa_indexer_topk, seq_length // 4) - avg_comp_4 = effective_topk_4 * (1 - effective_topk_4 * 4 / (2 * seq_length)) - sparse_attn_r4 = ( - n_layers_r4 * num_attention_heads * (csa_window_size + avg_comp_4) * v_head_dim * 2 - ) - # Indexer token-linear: compressor (coff=2, wkv + wgate), Q proj, weights proj. - indexer_token_term = ( - n_layers_r4 * hidden_size * (2 * dsa_indexer_head_dim) * 2 - + n_layers_r4 * q_lora_rank * dsa_indexer_n_heads * dsa_indexer_head_dim - + n_layers_r4 * hidden_size * dsa_indexer_n_heads - ) - # Indexer L^2: scoring each query against ~L/4 compressed positions. - indexer_scoring_core = n_layers_r4 * dsa_indexer_n_heads * dsa_indexer_head_dim / 4 - else: - sparse_attn_r4 = 0 - indexer_token_term = 0 - indexer_scoring_core = 0 - - token_linear = ( - mla_proj_term - + sparse_attn_r0 - + sparse_attn_r4 - + sparse_attn_r128_window - + main_compressor_term - + indexer_token_term - ) - core = sparse_attn_r128_core + indexer_scoring_core - return token_linear, core - - def num_floating_point_operations( args, batch_size, seqlen_squared_sum_in_batch=None, total_real_tokens_in_batch=None ): @@ -596,49 +489,6 @@ def attn_layer_flops( + 2 * seqlen_squared_sum * hidden_size * p ) - def mla_attn_layer_flops( - total_tokens, - seqlen_squared_sum, - hidden_size, - num_heads, - q_lora_rank, - kv_lora_rank, - qk_head_dim, - qk_pos_emb_head_dim, - v_head_dim, - ): - """Calculate FLOPs for a Multi-Latent Attention (MLA) layer. - - Mirrors the MLA term in ``transformer_flops`` so the hybrid estimate - matches the standard-model estimate per attention layer (the DSv4 - attention layers -- Window/CSA/HCA -- are all MLA-based). The generic - ``attn_layer_flops`` above assumes dense MHA/GQA QKV projections and - badly overcounts MLA, whose Q/KV go through low-rank ``lora`` ranks. - - Returns a forward-equivalent value (FMA factor 2 baked in, NO fwd+bwd - factor): the caller (``hybrid_flops``) applies the global ``* 3`` for - forward + wgrad + dgrad, exactly as for the other layer types. - """ - fma = 2 - if q_lora_rank is None: - q_term = hidden_size * num_heads * (qk_head_dim + qk_pos_emb_head_dim) - else: - q_term = q_lora_rank * ( - hidden_size + num_heads * (qk_head_dim + qk_pos_emb_head_dim) + 1 - ) - # Token-linear part (q lora+rope+norm, kv lora+rope+norm, output proj). - token_linear = fma * ( - q_term - + kv_lora_rank * (hidden_size + num_heads * (qk_head_dim + v_head_dim) + 1) - + hidden_size * qk_pos_emb_head_dim - + (num_heads * v_head_dim) * hidden_size - ) - # Core attention (L^2) part: QK^T and (softmax(QK^T))V. /2 (causal) cancels *2 (FMA). - core = fma * ( - num_heads * (qk_head_dim + qk_pos_emb_head_dim) / 2 + num_heads * v_head_dim / 2 - ) - return token_linear * total_tokens + core * seqlen_squared_sum - def mamba_layer_flops( total_tokens, hidden_size, state_dim=16, head_dim=64, num_groups=1, num_heads=128 ): @@ -715,65 +565,11 @@ def hybrid_flops( gdn_conv_kernel_dim=4, vocab_size=256000, mtp_num_layers=0, - multi_latent_attention=False, - q_lora_rank=None, - kv_lora_rank=0, - qk_head_dim=0, - qk_pos_emb_head_dim=0, - v_head_dim=0, - experimental_attention_variant=None, - dsv4_n_layers_r0=0, - dsv4_n_layers_r4=0, - dsv4_n_layers_r128=0, - o_groups=None, - o_lora_rank=None, - csa_window_size=None, - seq_length=None, - dsa_indexer_n_heads=None, - dsa_indexer_head_dim=None, - dsa_indexer_topk=None, ): """Calculate total FLOPs for the hybrid model.""" - # Self-attention (already summed over all attention layers, fwd-equivalent - # with the FMA factor baked in; the global ``* 3`` below adds fwd+bwd). - if experimental_attention_variant == "dsv4_hybrid": - # DSv4 attention (Window/CSA/HCA) is MLA-based but runs SPARSE - # attention, not full O(L^2) MLA. Use the shared helper -- the single - # source of truth also used by ``transformer_flops`` -- so both code - # paths agree. Window/CSA/HCA layer counts come from the pattern. - dsv4_token_term, dsv4_core_term = _dsv4_hybrid_self_attention_flops( - hidden_size=hidden_size, - num_attention_heads=num_attn_heads, - v_head_dim=v_head_dim, - q_lora_rank=q_lora_rank, - o_groups=o_groups, - o_lora_rank=o_lora_rank, - csa_window_size=csa_window_size, - seq_length=seq_length, - n_layers_r0=dsv4_n_layers_r0, - n_layers_r4=dsv4_n_layers_r4, - n_layers_r128=dsv4_n_layers_r128, - dsa_indexer_n_heads=dsa_indexer_n_heads, - dsa_indexer_head_dim=dsa_indexer_head_dim, - dsa_indexer_topk=dsa_indexer_topk, - ) - attn_flops_total = 2 * ( - dsv4_token_term * total_tokens + dsv4_core_term * seqlen_squared_sum - ) - elif multi_latent_attention: - attn_flops_total = num_attn_layers * mla_attn_layer_flops( - total_tokens, - seqlen_squared_sum, - hidden_size, - num_attn_heads, - q_lora_rank, - kv_lora_rank, - qk_head_dim, - qk_pos_emb_head_dim, - v_head_dim, - ) - else: - attn_flops_total = num_attn_layers * attn_layer_flops( + flops_fwd = ( + num_attn_layers + * attn_layer_flops( total_tokens, seqlen_squared_sum, hidden_size, @@ -782,8 +578,6 @@ def hybrid_flops( gqa_groups, kv_channels, ) - flops_fwd = ( - attn_flops_total + num_mlp_layers * mlp_layer_flops(total_tokens, hidden_size, mlp_expansion, swiglu) + num_mamba_layers * mamba_layer_flops( @@ -814,9 +608,6 @@ def hybrid_flops( gdn_num_v_heads, gdn_conv_kernel_dim, ) - + - # MTP norms (eh_norm + final_norm) and eh projection (2 * h^2). - 2 * mtp_num_layers * (3 * hidden_size + 2 * hidden_size * hidden_size) * total_tokens + ( 2 * total_tokens * hidden_size * vocab_size * (1 + mtp_num_layers) ) # logits computation @@ -908,57 +699,48 @@ def transformer_flops(): https://arxiv.org/abs/2305.10403 https://arxiv.org/abs/2205.05198 ''' - if args.experimental_attention_variant == "dsv4_hybrid": - ## DSv4 hybrid: the MLA projections AND the sparse attention that - ## replaces full core attention are both computed together by - ## ``_dsv4_hybrid_self_attention_flops`` in the dsv4_hybrid branch - ## below (the single source of truth shared with ``hybrid_flops``). - ## Zero the standard MLA terms here so they are not double-counted. - standard_self_attn_term = 0 - standard_self_attn_core_term = 0 + ## MLA + if args.q_lora_rank is None: + q_term = ( + args.hidden_size + * args.num_attention_heads + * (args.qk_head_dim + args.qk_pos_emb_head_dim) + ) else: - ## MLA - if args.q_lora_rank is None: - q_term = ( - args.hidden_size - * args.num_attention_heads - * (args.qk_head_dim + args.qk_pos_emb_head_dim) - ) - else: - q_term = args.q_lora_rank * ( + q_term = args.q_lora_rank * ( + args.hidden_size + + args.num_attention_heads * (args.qk_head_dim + args.qk_pos_emb_head_dim) + + 1 + ) + # Token-linear part of MLA self-attention (lora projs, kv proj, RoPE, output proj). + standard_self_attn_term = ( + forward_backward_expansion_factor + * fma_expansion_factor + * ( + ## q lora + rope + q norm + q_term + ## kv lora + rope + kv norm + + args.kv_lora_rank + * ( args.hidden_size - + args.num_attention_heads * (args.qk_head_dim + args.qk_pos_emb_head_dim) + + args.num_attention_heads * (args.qk_head_dim + args.v_head_dim) + 1 ) - # Token-linear part of MLA self-attention (lora projs, kv proj, RoPE, output proj). - standard_self_attn_term = ( - forward_backward_expansion_factor - * fma_expansion_factor - * ( - ## q lora + rope + q norm - q_term - ## kv lora + rope + kv norm - + args.kv_lora_rank - * ( - args.hidden_size - + args.num_attention_heads * (args.qk_head_dim + args.v_head_dim) - + 1 - ) - + args.hidden_size * args.qk_pos_emb_head_dim - ## o proj - + (args.num_attention_heads * args.v_head_dim) * args.hidden_size - ) + + args.hidden_size * args.qk_pos_emb_head_dim + ## o proj + + (args.num_attention_heads * args.v_head_dim) * args.hidden_size ) - # Core-attention (L^2) part: ``QK^T`` and ``(softmax(QK^T)) V``. The - # ``/2`` accounts for the causal mask and the ``*2`` cancels it via FMA. - standard_self_attn_core_term = ( - forward_backward_expansion_factor - * fma_expansion_factor - * ( - args.num_attention_heads * (args.qk_head_dim + args.qk_pos_emb_head_dim) / 2 - + args.num_attention_heads * args.v_head_dim / 2 - ) + ) + # Core-attention (L^2) part: ``QK^T`` and ``(softmax(QK^T)) V``. The + # ``/2`` accounts for the causal mask and the ``*2`` cancels it via FMA. + standard_self_attn_core_term = ( + forward_backward_expansion_factor + * fma_expansion_factor + * ( + args.num_attention_heads * (args.qk_head_dim + args.qk_pos_emb_head_dim) / 2 + + args.num_attention_heads * args.v_head_dim / 2 ) + ) else: ## MHA or GQA @@ -992,8 +774,6 @@ def transformer_flops(): * 2 # QK^T and (QK^T)V ) - dsv4_hybrid_extra_term = 0 - dsv4_hybrid_extra_core_term = 0 if is_linear_attention_variant(args.experimental_attention_variant): # Calculate number of dense and MoE Transformer MLPs. if isinstance(args.linear_attention_freq, int): @@ -1051,45 +831,6 @@ def transformer_flops(): "Invalid experimental_attention_variant: " f"{args.experimental_attention_variant}" ) - elif args.experimental_attention_variant == "dsv4_hybrid": - # DSv4 hybrid: MLA projections + sparse attention (window / - # compressed-KV), main compressor, and learned indexer (DSA) are all - # computed by the shared ``_dsv4_hybrid_self_attention_flops`` helper. - # The standard MLA terms were zeroed above to avoid double-counting. - num_linear_attention_layers = 0 - linear_self_attn_term = 0 - num_standard_attention_layers = num_layers - - compress_ratios = args.csa_compress_ratios - assert compress_ratios is not None, "csa_compress_ratios must be set for dsv4_hybrid" - assert len(compress_ratios) == num_layers, ( - f"Invalid length of csa_compress_ratios: {len(compress_ratios)}, " - f"expected num_layers + mtp_num_layers ({num_layers})." - ) - # ratio == 0: window-only; ratio == 4: window + topk compressed KV - # (compressor + indexer); ratio == 128: window + all compressed KV. - dsv4_token_term, dsv4_core_term = _dsv4_hybrid_self_attention_flops( - hidden_size=args.hidden_size, - num_attention_heads=args.num_attention_heads, - v_head_dim=args.v_head_dim, - q_lora_rank=args.q_lora_rank, - o_groups=args.o_groups, - o_lora_rank=args.o_lora_rank, - csa_window_size=args.csa_window_size, - seq_length=args.seq_length, - n_layers_r0=sum(1 for r in compress_ratios if r == 0), - n_layers_r4=sum(1 for r in compress_ratios if r == 4), - n_layers_r128=sum(1 for r in compress_ratios if r == 128), - dsa_indexer_n_heads=args.dsa_indexer_n_heads, - dsa_indexer_head_dim=args.dsa_indexer_head_dim, - dsa_indexer_topk=args.dsa_indexer_topk, - ) - dsv4_hybrid_extra_term = ( - forward_backward_expansion_factor * fma_expansion_factor * dsv4_token_term - ) - dsv4_hybrid_extra_core_term = ( - forward_backward_expansion_factor * fma_expansion_factor * dsv4_core_term - ) else: num_linear_attention_layers = 0 linear_self_attn_term = 0 @@ -1100,14 +841,9 @@ def transformer_flops(): self_attn_term = ( linear_self_attn_term * num_linear_attention_layers + standard_self_attn_term * num_standard_attention_layers - + dsv4_hybrid_extra_term - ) - # Core attention (L^2) FLOPs. Standard attention has a uniform per-layer - # coefficient; DSv4 sparse attention varies by layer type and is pre-summed. - self_attn_core_term = ( - standard_self_attn_core_term * num_standard_attention_layers - + dsv4_hybrid_extra_core_term ) + # Core attention (L^2) FLOPs per standard-attention layer. + self_attn_core_term = standard_self_attn_core_term * num_standard_attention_layers # Token-linear FLOPs scale with the real (unpadded) token count. # For BSHD this falls back to ``batch_size * seq_length`` (no padding). @@ -1177,31 +913,11 @@ def transformer_flops(): get_hybrid_layer_counts, ) - layer_counts = get_hybrid_layer_counts(args.hybrid_layer_pattern) - num_mamba_layers, num_gdn_layers, num_mlp_layers, num_moe_layers = itemgetter( - Symbols.MAMBA, Symbols.GDN, Symbols.MLP, Symbols.MOE - )(layer_counts) - # Attention layers = plain ATTENTION ('*') PLUS every MLA variant - # (DS_ATTENTION 'D', CSA 'C', HCA 'H', WINDOW 'W'). Previously only '*' - # was counted, so a DSv4 pattern (all C/H/W, zero '*') yielded - # num_attn_layers=0 and dropped ALL attention FLOPs from the estimate, - # roughly halving the reported throughput vs. the gpt_model. - num_attn_layers = layer_counts[Symbols.ATTENTION] + sum( - layer_counts[s] for s in Symbols.MLA_ATTENTION - ) - - # DSv4 sparse-attention layer counts by compression ratio, derived from the - # pattern symbols: WINDOW -> r0, CSA -> r4, HCA -> r128 (mirrors the - # ``csa_compress_ratios`` 0/4/128 buckets used by ``transformer_flops``). - dsv4_n_layers_r0 = layer_counts[Symbols.WINDOW] - dsv4_n_layers_r4 = layer_counts[Symbols.CSA] - dsv4_n_layers_r128 = layer_counts[Symbols.HCA] - if args.experimental_attention_variant == "dsv4_hybrid": - assert num_attn_layers == (dsv4_n_layers_r0 + dsv4_n_layers_r4 + dsv4_n_layers_r128), ( - "dsv4_hybrid expects all attention layers to be Window/CSA/HCA; " - f"got {num_attn_layers} attention layers but only " - f"{dsv4_n_layers_r0 + dsv4_n_layers_r4 + dsv4_n_layers_r128} are W/C/H." + num_mamba_layers, num_gdn_layers, num_attn_layers, num_mlp_layers, num_moe_layers = ( + itemgetter(Symbols.MAMBA, Symbols.GDN, Symbols.ATTENTION, Symbols.MLP, Symbols.MOE)( + get_hybrid_layer_counts(args.hybrid_layer_pattern) ) + ) mtp_num_layers = args.mtp_num_layers if mtp_num_layers is None: @@ -1245,23 +961,6 @@ def transformer_flops(): gdn_conv_kernel_dim=args.linear_conv_kernel_dim or 4, vocab_size=args.padded_vocab_size, mtp_num_layers=mtp_num_layers, - multi_latent_attention=args.multi_latent_attention, - q_lora_rank=args.q_lora_rank, - kv_lora_rank=args.kv_lora_rank, - qk_head_dim=args.qk_head_dim, - qk_pos_emb_head_dim=args.qk_pos_emb_head_dim, - v_head_dim=args.v_head_dim, - experimental_attention_variant=args.experimental_attention_variant, - dsv4_n_layers_r0=dsv4_n_layers_r0, - dsv4_n_layers_r4=dsv4_n_layers_r4, - dsv4_n_layers_r128=dsv4_n_layers_r128, - o_groups=getattr(args, "o_groups", None), - o_lora_rank=getattr(args, "o_lora_rank", None), - csa_window_size=getattr(args, "csa_window_size", None), - seq_length=args.seq_length, - dsa_indexer_n_heads=getattr(args, "dsa_indexer_n_heads", None), - dsa_indexer_head_dim=getattr(args, "dsa_indexer_head_dim", None), - dsa_indexer_topk=getattr(args, "dsa_indexer_topk", None), ) else: # Compute standard Transformer model FLOPs. @@ -1290,7 +989,7 @@ def get_start_time_from_progress_log(): def _get_field(string, type): return type(string.split(': ')[1]) - with open_file(progress_log_filename, 'r') as f: + with maybe_msc.open(progress_log_filename, 'r') as f: for line in f: line = line.strip() line_tokens = line.split('\t') @@ -1359,7 +1058,15 @@ def reorder_inner_param_groups(optimizer_state_dict): if "param_groups" not in inner_optimizer: return param_groups = inner_optimizer["param_groups"] - key_fn = lambda pg: [pg[key] for key in param_group_identifier_keys] + + # Treat missing and explicit None identifier values as equivalent. + # Wrap each component so None never compares directly with floats or strings. + def key_fn(pg): + return [ + (value is not None, value) + for value in (pg.get(key) for key in param_group_identifier_keys) + ] + param_groups.sort(key=key_fn) inner_optimizer["param_groups"] = param_groups @@ -1973,6 +1680,7 @@ def wrap_model_chunks_with_ddp( ddp_config, *, use_layer_wise_distributed_optimizer=False, + use_layer_wise_param_layout=True, DP=DDP, pg_collection=None, bucket_sizes=None, @@ -1983,13 +1691,12 @@ def wrap_model_chunks_with_ddp( Centralises the DDP-wrapping wiring shared between :func:`get_model` and unit tests. - For ``use_layer_wise_distributed_optimizer=True``: forces - ``ddp_config.use_distributed_optimizer=True`` (mutated in place; needed for reduce-scatter) - and computes per-chunk layouts via - :meth:`LayerWiseDistributedOptimizer.compute_full_param_layout`. That method picks the - padded shard-aligned LayerWise layout or the compact decoupled layout per - ``ddp_config.use_layer_wise_param_layout`` (True → padded, False → compact, with LayerWise - buffers treated as non-DistOpt and synced via legacy ``allgather_params``). + For ``use_layer_wise_distributed_optimizer=True`` and ``use_layer_wise_param_layout=True``: + forces ``ddp_config.use_distributed_optimizer=True`` (mutated in place; needed + for reduce-scatter), and computes per-chunk shard-aligned layouts via + :meth:`LayerWiseDistributedOptimizer.compute_full_param_layout`. With + ``use_layer_wise_param_layout=False``, no layout is supplied and LayerWise falls back + to its legacy ``allgather_params`` sync path. For non-layerwise with ``ddp_config.use_distributed_optimizer=True``: computes per-chunk byte-level layouts via @@ -2005,9 +1712,11 @@ def wrap_model_chunks_with_ddp( model_chunks: List of model chunks to wrap (un-DDP-wrapped). config: :class:`TransformerConfig`. ddp_config: :class:`DistributedDataParallelConfig`. Mutated in place when - ``use_layer_wise_distributed_optimizer=True``. Its ``use_layer_wise_param_layout`` - field selects the padded vs compact decoupled LayerWise layout (default). + ``use_layer_wise_distributed_optimizer=True`` and ``use_layer_wise_param_layout=True``. use_layer_wise_distributed_optimizer: Whether the layerwise wiring runs. + use_layer_wise_param_layout: When ``use_layer_wise_distributed_optimizer=True``, + controls whether to compute and supply a shard-aligned param layout + to DDP. ``False`` keeps LayerWise on its legacy sync path. DP: The DDP class to construct (``DistributedDataParallel`` or an FSDP variant). pg_collection: Optional :class:`ProcessGroupCollection`. When provided, @@ -2030,18 +1739,13 @@ def wrap_model_chunks_with_ddp( # Compute per-chunk layouts (DDP only). per_chunk_layouts = [None] * n if DP is DDP: - if use_layer_wise_distributed_optimizer: - # LayerWise (Muon) optimizer. Force use_distributed_optimizer=True so sibling - # non-LayerWise buffers (embeddings, biases, layernorm) shard with the byte-level - # DistributedOptimizer layout, and tag params so DDP buffer grouping routes - # LayerWise-managed matrices (Muon's Newton-Schulz domain) to a separate buffer. - # The padded-vs-compact LayerWise layout decision is made inside - # compute_full_param_layout / _ParamAndGradBuffer from use_layer_wise_param_layout: - # by default (compact) LayerWise buffers get the no-padding layout and the per-buffer - # override flips use_distributed_optimizer off for them; with - # --use-layer-wise-param-layout they stay on the padded DistOpt layout. + if use_layer_wise_distributed_optimizer and use_layer_wise_param_layout: ddp_config.use_distributed_optimizer = True compute_layout = LayerWiseDistributedOptimizer.compute_full_param_layout + # Tag params so DDP buffer grouping routes LayerWise-managed matrices + # (Muon's Newton-Schulz domain) to a shard-aligned buffer and routes + # everything else (embeddings, biases, layernorm) to a separate + # DistOpt-style buffer. tag_params_for_buffer_routing(model_chunks) elif not use_layer_wise_distributed_optimizer and ddp_config.use_distributed_optimizer: compute_layout = DistributedOptimizer.compute_full_param_layout @@ -2057,8 +1761,22 @@ def wrap_model_chunks_with_ddp( "wrap_model_chunks_with_ddp requires a dp_cp process group to size " "the distributed-optimizer parameter layout" ) - data_parallel_world_size = get_pg_size(layout_pgs.dp_cp) - expert_data_parallel_world_size = get_pg_size(getattr(layout_pgs, "expt_dp", None)) + # The distributed optimizer shards each bucket over the intra-instance group, which + # is what DDP hands to the buffer as its data_parallel_group. Size the layout by that + # same group, otherwise the layout reports more shards than the reduce-scatter uses + # and the trailing shards of every bucket end up owned by no rank. intra_dp_cp is + # the full dp_cp when num_distributed_optimizer_instances is 1, so this only differs + # when there are several instances. + intra_dp_cp_group = getattr(layout_pgs, "intra_dp_cp", None) + intra_expt_dp_group = getattr(layout_pgs, "intra_expt_dp", None) + data_parallel_world_size = get_pg_size( + intra_dp_cp_group if intra_dp_cp_group is not None else layout_pgs.dp_cp + ) + expert_data_parallel_world_size = get_pg_size( + intra_expt_dp_group + if intra_expt_dp_group is not None + else getattr(layout_pgs, "expt_dp", None) + ) for i, (chunk, bucket_size) in enumerate(zip(model_chunks, bucket_sizes)): all_params = [p for p in chunk.parameters() if p.requires_grad] per_chunk_layouts[i] = compute_layout( @@ -2092,6 +1810,31 @@ def wrap_model_chunks_with_ddp( return wrapped +def _freeze_all_model_chunks(model_list): + """Freeze all parameters in a list of model chunks (for logits-saving runs).""" + for model_module in model_list: + model_module.requires_grad_(False) + # Additionally freeze expert biases of routers + for module in model_module.modules(): + if hasattr(module, "frozen_expert_bias"): + module.frozen_expert_bias = True + return model_list + + +def _forward_backward_grad_context(args): + """Grad context for a train step's forward/backward pass. + + Returns a tuple of (grad_context, forward_only). + grad_context is ``torch.no_grad()`` when all layers are frozen (e.g. teacher logits + dumps), no parameter needs gradients, so there is no reason to build the + autograd graph. Otherwise returns a no-op context. + forward_only is True when all layers are frozen, False otherwise. + """ + grad_context = torch.no_grad() if getattr(args, "freeze_all_layers", False) else nullcontext() + forward_only = getattr(args, "freeze_all_layers", False) + return grad_context, forward_only + + def get_model( model_provider_func, model_type=ModelType.encoder_or_decoder, @@ -2122,16 +1865,7 @@ def get_model( print_rank_0("> including expert parallelism AG group") if has_nvidia_modelopt: - from megatron.post_training.checkpointing import has_modelopt_state - - # [ModelOpt]: Check if the checkpoint is a ModelOpt checkpoint and - # set a flag to use our model provider if so. - if args.load is not None and has_modelopt_state(args.load): - print_rank_0(f'ModelOpt checkpoint detected') - args.modelopt_enabled = True - elif getattr(args, "export_kd_teacher_load", None): - # For distillation ckpts without ModelOpt state - args.modelopt_enabled = True + maybe_enable_modelopt(args) # Build model. def build_model(): @@ -2182,12 +1916,7 @@ def build_model(): # For rare operations like post-training logits saving if args.freeze_all_layers: - for model_module in model: - model_module.requires_grad_(False) - # Additionally freeze expert biases of routers - for module in model_module.modules(): - if hasattr(module, "frozen_expert_bias"): - module.frozen_expert_bias = True + _freeze_all_model_chunks(model) # Set tensor model parallel attributes if not set. # Only parameters that are already tensor model parallel have these @@ -2203,9 +1932,12 @@ def build_model(): ) if get_pg_rank(pg_collection.dp) == 0 and get_pg_rank(pg_collection.cp) == 0: print( - ' > number of parameters on (tensor, pipeline) ' - 'model parallel rank ({}, {}): {}'.format( - get_pg_rank(pg_collection.tp), get_pg_rank(pg_collection.pp), num_parameters + ' > number of parameters on (tensor, gtp_weight_remat, pipeline) ' + 'model parallel rank ({}, {}, {}): {}'.format( + get_pg_rank(pg_collection.tp), + get_pg_rank(pg_collection.gtp_remat), + get_pg_rank(pg_collection.pp), + num_parameters, ), flush=True, ) @@ -2270,12 +2002,15 @@ def build_model(): for disable in per_chunk_disable_bucketing ] - # Setup stream for ddp initialization. The side-stream may be necessary for cuda graph - # capture support with DDP, but we sync it with the current stream to avoid races. - ddp_stream = torch.cuda.Stream() - # Wait for the default stream to complete before starting ddp_stream + if config.cuda_graph_impl == "full_iteration": + # DDP initialization must use the full-iteration capture stream so its retained + # AccumulateGrad nodes do not reference a different, non-capturing stream. + ddp_stream = get_shared_capture_stream() + else: + # Preserve a dedicated initialization stream for all other implementations. + ddp_stream = torch.cuda.Stream() ddp_stream.wait_stream(torch.cuda.current_stream()) - # Make ddp_stream start after whatever the default stream already queued + with torch.cuda.stream(ddp_stream): model = wrap_model_chunks_with_ddp( model, @@ -2284,13 +2019,13 @@ def build_model(): use_layer_wise_distributed_optimizer=getattr( args, 'use_layer_wise_distributed_optimizer', False ), + use_layer_wise_param_layout=getattr(args, 'use_layer_wise_param_layout', True), DP=DP, pg_collection=pg_collection if args.use_megatron_fsdp else None, bucket_sizes=per_chunk_bucket_sizes, disable_bucketing_per_chunk=per_chunk_disable_bucketing, ) - # End of setup_stream - # Critical: ensure side-stream work completes before touching params on default stream + # Ensure initialization-stream work completes before touching params on the default stream. torch.cuda.current_stream().wait_stream(ddp_stream) # Broadcast params from data parallel src rank to other data parallel ranks. @@ -2436,6 +2171,9 @@ def setup_model_and_optimizer( skip_optimizer = not (has_normal_optimizer or has_rl_optimizer) wrap_with_ddp = not skip_optimizer + if has_nvidia_modelopt: + maybe_enable_modelopt(args) + def _build_model_wrapper(wrap_with_ddp: bool): if cfg_container is not None and getattr(cfg_container, "model", None) is not None: from megatron.training.utils import start_memory_history_recording @@ -2446,6 +2184,12 @@ def _build_model_wrapper(wrap_with_ddp: bool): model_config = cfg.model builder_cls = model_config.get_builder_cls() builder = builder_cls(model_config) + + # Inject freeze_all_layers as a pre-wrap hook so DDP sees requires_grad=False + # and skips grad-buffer allocation for all params (matching get_model behavior). + if args.freeze_all_layers: + model_config.pre_wrap_hooks.append(_freeze_all_model_chunks) + return builder.build_distributed_models( pg_collection=pg_collection, ddp_config=cfg.ddp, @@ -2454,7 +2198,6 @@ def _build_model_wrapper(wrap_with_ddp: bool): use_torch_fsdp2=cfg.dist.use_torch_fsdp2, wrap_with_ddp=wrap_with_ddp, data_parallel_random_init=cfg.rng.data_parallel_random_init, - use_layer_wise_distributed_optimizer=cfg.optimizer.use_layer_wise_distributed_optimizer, ) else: assert ( @@ -2467,10 +2210,36 @@ def _build_model_wrapper(wrap_with_ddp: bool): pg_collection=pg_collection, ) + # Configure GTP weight-remat padding/loss reduction before model construction (pad + # alignment governs how dim-0 shards are built). Placed here (not in get_model) so it + # also covers the config-container builder path, which does not call get_model. + if is_gtp_remat_active(args): + from megatron.core.tensor_parallel.gtp_api import configure_gtp_remat_from_recipe + + configure_gtp_remat_from_recipe( + fp4=getattr(args, 'fp4', None) is not None, + fp8_recipe=getattr(args, 'fp8_recipe', None), + fp8=getattr(args, 'fp8', None) is not None, + calculate_per_token_loss=getattr(args, 'calculate_per_token_loss', False), + ) + model = _build_model_wrapper(wrap_with_ddp) unwrapped_model = unwrap_model(model) - if args.logits_save_dir is not None: + # Classify each GTP param's prefetch chain after model build + DDP wrap, before the + # first forward. Placed here (not in get_model) so it also covers the config-container + # builder path. + if is_gtp_remat_active(args): + from megatron.core.tensor_parallel.gtp_api import classify_gtp_remat_chains + + classify_gtp_remat_chains( + model, + cuda_graph_modules=getattr(args, 'cuda_graph_modules', None), + moe_shared_expert_overlap=getattr(args, 'moe_shared_expert_overlap', False), + cuda_graph_impl=getattr(args, 'cuda_graph_impl', 'none'), + ) + + if args.logits_save_dir is not None and mpu.is_pipeline_last_stage(): from megatron.training.distillation import LogitsSaverHooks logits_saver = LogitsSaverHooks( @@ -2482,7 +2251,7 @@ def _build_model_wrapper(wrap_with_ddp: bool): ) logits_saver.attach_hooks(unwrapped_model[-1]) - if args.logits_load_dir is not None: + if args.logits_load_dir is not None and mpu.is_pipeline_last_stage(): from megatron.training.distillation import StudentLogitsCapture student_logits_capture = StudentLogitsCapture() @@ -2591,7 +2360,9 @@ def _build_model_wrapper(wrap_with_ddp: bool): and args.ckpt_format == "torch_dist", tp_group=ckpt_pgc.tp if ckpt_pgc is not None else None, pp_group=ckpt_pgc.pp if ckpt_pgc is not None else None, - dp_cp_group=ckpt_pgc.dp_cp if ckpt_pgc is not None else None, + # Replica_id must match the save path (see save_checkpoint_and_time): use the + # gtp_remat-inclusive group, not replicate dp_cp, or gtp_remat peers collide. + dp_cp_group=getattr(ckpt_pgc, "dp_cp_gtp_remat", None), dp_group=ckpt_pgc.dp if ckpt_pgc is not None else None, expt_dp_group=ckpt_pgc.expt_dp if ckpt_pgc is not None else None, rng_state_key_prefix=getattr(unwrapped_model[0], "rng_state_key_prefix", ""), @@ -2608,6 +2379,14 @@ def _build_model_wrapper(wrap_with_ddp: bool): args.iteration = 0 args.num_floating_point_operations_so_far = 0 + # [ModelOpt]: Load the teacher checkpoint for ModelOpt distillation if applicable. + # Import locally to prevent circular import: megatron.post_training.checkpointing + # imports `get_args` from megatron.training at module scope. + if has_nvidia_modelopt: + from megatron.post_training.checkpointing import load_kd_teacher_checkpoint + + load_kd_teacher_checkpoint(model) + # Validate that the world size can accommodate the current batch size. # This catches the case where GPUs were scaled up mid-training but the # current position in the batch size schedule yields a batch size that @@ -2742,37 +2521,8 @@ def train_step( """ args = get_args() timers = get_timers() - num_microbatches = get_num_microbatches() - - offload_optimizer_states = getattr(args, 'offload_optimizer_states', False) - if offload_optimizer_states: - # Reload optimizer states as late as possible so the H2D transfer can overlap - # with gradient finalization. Preserve custom finalize hooks installed by a - # model builder, and avoid wrapping the hook again on every training step. - finalize_model_grads_func = getattr(config, 'finalize_model_grads_func', None) - if ( - getattr(finalize_model_grads_func, '_optimizer_state_offload_wrapped_optimizer', None) - is not optimizer - ): - base_finalize_model_grads_func = finalize_model_grads_func or finalize_model_grads - - def finalize_model_grads_with_state_reload(*fmg_args, **fmg_kwargs): - for optim_instance in optimizer.chained_optimizers: - if isinstance(optim_instance, DistributedOptimizer): - optim_instance.reload_offloaded_states() - return base_finalize_model_grads_func(*fmg_args, **fmg_kwargs) - - setattr( - finalize_model_grads_with_state_reload, - '_optimizer_state_offload_wrapped_optimizer', - optimizer, - ) - config.finalize_model_grads_func = finalize_model_grads_with_state_reload rerun_state_machine = get_rerun_state_machine() - packed_data_iterator = None - has_wrapped_data_iterator = False - rerun_data_iterator = data_iterator save_params_in_this_iteration = ( args.save_params_interval is not None and (iteration + 1) % args.save_params_interval == 0 ) @@ -2790,13 +2540,7 @@ def finalize_model_grads_with_state_reload(*fmg_args, **fmg_kwargs): save_dgrads_in_this_iteration = ( args.save_dgrads_interval is not None and (iteration + 1) % args.save_dgrads_interval == 0 ) - while rerun_state_machine.should_run_forward_backward(rerun_data_iterator): - # Start the D2H transfer before zeroing gradients to maximize overlap. - if offload_optimizer_states: - for optim_instance in optimizer.chained_optimizers: - if isinstance(optim_instance, DistributedOptimizer): - optim_instance.offload_states() - + while rerun_state_machine.should_run_forward_backward(data_iterator): # Set grad to zero. for model_chunk in model: model_chunk.zero_grad_buffer() @@ -2838,13 +2582,6 @@ def finalize_model_grads_with_state_reload(*fmg_args, **fmg_kwargs): if isinstance(optim_instance, DistributedOptimizer): optim_instance._copy_main_params_to_param_buffer() - # Master weights must remain resident until any main-param copy above is - # complete. Releasing here keeps optimizer memory out of forward/backward. - if offload_optimizer_states: - for optim_instance in optimizer.chained_optimizers: - if isinstance(optim_instance, DistributedOptimizer): - optim_instance.release_offloaded_gpu_states() - # Forward pass. if save_activations_in_this_iteration: enable_activation_logging(model, args.save) @@ -2852,39 +2589,22 @@ def finalize_model_grads_with_state_reload(*fmg_args, **fmg_kwargs): enable_tokens_per_expert_logging(model, args.save) if save_dgrads_in_this_iteration: enable_dgrad_logging(model, args.save) - if getattr(config, 'sequence_packing_scheduler', None) is not None: - # Dynamic-CP / sequence packing must happen after the rerun state machine has - # observed the original RerunDataIterator. The scheduler returns another - # RerunDataIterator containing the packed microbatches for this step. - if not has_wrapped_data_iterator: - ( - packed_data_iterator, - num_microbatches, - seqlen_sum_this_global_batch, - seqlen_squared_sum_this_global_batch, - ) = wrap_data_iterator(data_iterator, config, get_num_microbatches()) - has_wrapped_data_iterator = True - rerun_data_iterator = packed_data_iterator - forward_backward_data_iterator = packed_data_iterator - else: - num_microbatches = get_num_microbatches() - seqlen_sum_this_global_batch = args.seq_length * args.global_batch_size - seqlen_squared_sum_this_global_batch = args.seq_length**2 * args.global_batch_size - forward_backward_data_iterator = data_iterator - losses_reduced = forward_backward_func( - forward_step_func=forward_step_func, - data_iterator=forward_backward_data_iterator, - model=model, - num_microbatches=num_microbatches, - seq_length=args.seq_length, - micro_batch_size=args.micro_batch_size, - decoder_seq_length=args.decoder_seq_length, - forward_only=False, - adjust_tensor_shapes_fn=adjust_tensor_shapes_fn, - force_all_reduce=save_wgrads_in_this_iteration, - p2p_communicator=p2p_communicator, - pg_collection=pg_collection, - ) + grad_context, forward_only = _forward_backward_grad_context(args) + with grad_context: + losses_reduced = forward_backward_func( + forward_step_func=forward_step_func, + data_iterator=data_iterator, + model=model, + num_microbatches=get_num_microbatches(), + seq_length=args.seq_length, + micro_batch_size=args.micro_batch_size, + decoder_seq_length=args.decoder_seq_length, + forward_only=forward_only, + adjust_tensor_shapes_fn=adjust_tensor_shapes_fn, + force_all_reduce=save_wgrads_in_this_iteration, + p2p_communicator=p2p_communicator, + pg_collection=pg_collection, + ) if save_activations_in_this_iteration: save_activations(iteration + 1) disable_activation_logging() @@ -2924,19 +2644,7 @@ def _save_state_dict(attr_name, label): should_checkpoint, should_exit, exit_code = rerun_state_machine.should_checkpoint_and_exit() if should_exit: - return ( - {}, - True, - should_checkpoint, - should_exit, - exit_code, - None, - None, - 0, - num_microbatches, - seqlen_sum_this_global_batch, - seqlen_squared_sum_this_global_batch, - ) + return {}, True, should_checkpoint, should_exit, exit_code, None, None, 0 # Empty unused memory. if args.empty_unused_memory_level >= 1: @@ -2973,7 +2681,9 @@ def _save_state_dict(attr_name, label): getattr(pg_collection, _required, None) is not None ), f"model pg_collection used by train_step must define {_required}" mp_group = pg_collection.mp - dp_cp_group = pg_collection.dp_cp + # gtp_remat-inclusive: the reported global per-token loss must cover gtp_remat peers' distinct + # tokens (replicate dp_cp would report a 1/gtp_remat subsample -> per-step noisy). Display-only. + dp_cp_group = getattr(pg_collection, 'dp_cp_gtp_remat', None) or pg_collection.dp_cp is_last_stage = is_pp_last_stage(pg_collection.pp) # when freezing sub-models we may have a mixture of successful and unsucessful ranks, # so we must gather across mp ranks @@ -2993,7 +2703,15 @@ def _save_state_dict(attr_name, label): # Update learning rate. if update_successful: - increment = get_num_microbatches() * args.micro_batch_size * args.data_parallel_size + # data_parallel_size excludes the GTP-remat axis (it's folded into total_model_size at + # arguments.py:446); each gtp-remat peer consumes a distinct microbatch, so multiply it + # back in for the sample count. + increment = ( + get_num_microbatches() + * args.micro_batch_size + * args.data_parallel_size + * args.gtp_weight_remat_size + ) opt_param_scheduler.step(increment=increment) skipped_iter = 0 else: @@ -3030,9 +2748,6 @@ def _save_state_dict(attr_name, label): grad_norm, num_zeros_in_grad, log_max_attention_logit, - num_microbatches, - seqlen_sum_this_global_batch, - seqlen_squared_sum_this_global_batch, ) return ( {}, @@ -3043,9 +2758,6 @@ def _save_state_dict(attr_name, label): grad_norm, num_zeros_in_grad, log_max_attention_logit, - num_microbatches, - seqlen_sum_this_global_batch, - seqlen_squared_sum_this_global_batch, ) @@ -3065,7 +2777,6 @@ def training_log( is_first_iteration=False, seqlen_squared_sum_in_batch: float | None = None, total_real_tokens_in_batch: float | None = None, - num_microbatches: int | None = None, ): """Log training information such as losses, timing, ....""" args = get_args() @@ -3144,8 +2855,15 @@ def training_log( if args.perform_rl_step: timers_to_log.extend(RL_LOGGABLE_TIMER_NAMES) - # Calculate batch size. - batch_size = args.micro_batch_size * args.data_parallel_size * get_num_microbatches() + # Calculate batch size. data_parallel_size excludes the GTP-remat axis (it's folded into + # total_model_size at arguments.py:446); each gtp-remat peer consumes a distinct microbatch, + # so multiply it back in for the global sample count. + batch_size = ( + args.micro_batch_size + * args.data_parallel_size + * args.gtp_weight_remat_size + * get_num_microbatches() + ) # Track app tag & app tag ID one_logger_utils.track_app_tag(batch_size, args.world_size, args.seq_length) @@ -3245,7 +2963,7 @@ def training_log( # Log MoE metrics. moe_log_string = "" if args.num_experts is not None: - moe_loss_scale = 1 / (num_microbatches or get_num_microbatches()) + moe_loss_scale = 1 / get_num_microbatches() track_names = [] if "aux_loss" in args.moe_router_load_balancing_type: track_names.append("load_balancing_loss") @@ -3256,23 +2974,15 @@ def training_log( if args.moe_z_loss_coeff is not None: track_names.append("z_loss") - moe_layer_freq = args.moe_layer_freq - mtp_num_layers = args.mtp_num_layers if is_hybrid_model(args): + from operator import itemgetter + from megatron.core.ssm.mamba_hybrid_layer_allocation import ( Symbols, - parse_hybrid_pattern, + get_hybrid_layer_counts, ) - parsed_hybrid_pattern = parse_hybrid_pattern(args.hybrid_layer_pattern) - main_pattern = (parsed_hybrid_pattern.main_pattern or "").replace(Symbols.PIPE, "") - layers = len(main_pattern) + parsed_hybrid_pattern.mtp_num_depths - moe_layer_freq = [int(layer_type == Symbols.MOE) for layer_type in main_pattern] - moe_layer_freq.extend( - parsed_hybrid_pattern.mtp_pattern.count(Symbols.MOE) - for _ in range(parsed_hybrid_pattern.mtp_num_depths) - ) - mtp_num_layers = None + layers = itemgetter(Symbols.MOE)(get_hybrid_layer_counts(args.hybrid_layer_pattern)) else: layers = args.num_layers @@ -3285,22 +2995,15 @@ def training_log( force_initialize=True, track_names=track_names, num_layers=layers, - moe_layer_freq=moe_layer_freq, - mtp_num_layers=mtp_num_layers, + moe_layer_freq=args.moe_layer_freq, + mtp_num_layers=args.mtp_num_layers, pg_collection=pg_collection, total_loss_dict=total_loss_dict, ) # Log MTP metrics. if args.mtp_num_layers is not None: - if args.calculate_per_token_loss: - # The tracker already reduces raw loss sums and token counts into a - # per-token loss, matching the main loss normalization path. - mtp_loss_scale = 1.0 - else: - # Legacy mode accumulates microbatch-normalized losses, so average - # by the scheduled microbatch count for this step. - mtp_loss_scale = 1 / (num_microbatches or get_num_microbatches()) + mtp_loss_scale = 1 / get_num_microbatches() MTPLossLoggingHelper.track_mtp_metrics( mtp_loss_scale, iteration, writer, wandb_writer, total_loss_dict ) @@ -3314,13 +3017,6 @@ def training_log( writer=writer, wandb_writer=wandb_writer, total_loss_dict=total_loss_dict, - num_layers=args.num_layers + (args.mtp_num_layers or 0), - num_indexer_layers=( - sum(ratio == 4 for ratio in args.csa_compress_ratios) - if args.csa_compress_ratios is not None - else None - ), - preserve_groups=args.cuda_graph_impl != "none", ) # Dump memory snapshot and print metrics to stdout. @@ -3505,9 +3201,25 @@ def compute_throughputs_and_append_to_progress_log(iteration, num_floating_point ) +def _assert_param_gather_overlap_model(model_chunk): + """Assert that a model chunk implements the parameter-gather overlap lifecycle.""" + # MimoModel is a composite wrapper rather than a DDP instance, but delegates this + # interface to its active inner DDP modules. + required_methods = ('enable_forward_pre_hook', 'disable_forward_pre_hook', 'start_param_sync') + missing_methods = [ + method_name + for method_name in required_methods + if not callable(getattr(model_chunk, method_name, None)) + ] + assert not missing_methods, ( + f'{type(model_chunk).__name__} does not support parameter-gather overlap; ' + f'missing callable methods: {", ".join(missing_methods)}' + ) + + def enable_forward_pre_hook(model_chunks): for model_chunk in model_chunks: - assert isinstance(model_chunk, DDP) + _assert_param_gather_overlap_model(model_chunk) model_chunk.enable_forward_pre_hook() @@ -3515,15 +3227,15 @@ def disable_forward_pre_hook(model_chunks, optimizer=None, param_sync=True): if param_sync and optimizer is not None: optimizer.prepare_model_params_for_param_sync() for model_chunk in model_chunks: - assert isinstance(model_chunk, DDP) + _assert_param_gather_overlap_model(model_chunk) model_chunk.disable_forward_pre_hook(param_sync=param_sync) -def force_param_sync(model_chunks: list[DDP], optimizer=None) -> None: +def force_param_sync(model_chunks, optimizer=None) -> None: if optimizer is not None: optimizer.prepare_model_params_for_param_sync() for model_chunk in model_chunks: - assert isinstance(model_chunk, DDP) + _assert_param_gather_overlap_model(model_chunk) model_chunk.start_param_sync(force_sync=True) @@ -3569,7 +3281,8 @@ def save_checkpoint_and_time( tp_group = getattr(ckpt_pgc, "tp", None) if ckpt_pgc is not None else None pp_group = getattr(ckpt_pgc, "pp", None) if ckpt_pgc is not None else None dp_group = getattr(ckpt_pgc, "dp", None) if ckpt_pgc is not None else None - dp_cp_group = getattr(ckpt_pgc, "dp_cp", None) if ckpt_pgc is not None else None + # Replica_id needs the gtp_remat-inclusive group (dp_cp_gtp_remat), not replicate dp_cp. + dp_cp_group = getattr(ckpt_pgc, "dp_cp_gtp_remat", None) if ckpt_pgc is not None else None expt_dp_group = getattr(ckpt_pgc, "expt_dp", None) if ckpt_pgc is not None else None # Per-grid rng key namespace set by a multi-grid model; '' for stock single-grid. rng_state_key_prefix = getattr(unwrap_model(model)[0], "rng_state_key_prefix", "") @@ -3931,12 +3644,15 @@ def train( ) def _dp_world_size(): + # Full DP x gtp_remat degree (num_microbatches spans the full data-distribution axis). + gtp_remat = args.gtp_weight_remat_size if lang_pgc is not None: - return lang_pgc.dp.size() + return lang_pgc.dp.size() * gtp_remat if mpu.model_parallel_is_initialized(): return mpu.get_data_parallel_world_size() - # args.data_parallel_size equals the language (llm) dp on all ranks (entry validate_args). - return args.data_parallel_size + # args.data_parallel_size is the language (llm) dp on all ranks (set in validate_args) and + # excludes gtp_remat, so scale by gtp_remat to span the full data-distribution axis. + return args.data_parallel_size * gtp_remat # IMPORTANT FIX: For RL training, reinitialize the microbatch calculator with the correct configuration if args.perform_rl_step: @@ -3963,6 +3679,9 @@ def _dp_world_size(): energy_monitor = get_energy_monitor() one_logger = get_one_logger() + if args.dynamic_context_parallel: + train_data_iterator = iter(HybridCPDataLoaderWrapper(train_data_iterator, config)) + if args.run_workload_inspector_server: try: import threading @@ -4202,7 +3921,6 @@ def trace_handler(p): seq_length=args.seq_length, micro_batch_size=args.micro_batch_size, optimizers=[optimizer], - thd_sequence_length_upper_bound=_get_thd_sequence_length_upper_bound(args), ) # Run training iterations till done. @@ -4246,9 +3964,9 @@ def trace_handler(p): # Standard microbatch update (sequence packing overrides this in rl_utils.py) update_num_microbatches(args.consumed_train_samples, consistency_check=False, verbose=True) # Skip automatic checkpoint on microbatch changes when sequence packing is active - # as it intentionally reconfigures microbatches. + # as it intentionally reconfigures microbatches if get_num_microbatches() != num_microbatches and iteration != 0: - if args.rl_use_sequence_packing or args.sequence_packing_scheduler is not None: + if args.rl_use_sequence_packing: print_rank_0( f"[Sequence Packing] Skipping automatic checkpoint at iteration {iteration} " f"(microbatch change: {num_microbatches} -> {get_num_microbatches()})" @@ -4334,9 +4052,6 @@ def trace_handler(p): grad_norm = 0.0 num_zeros_in_grad = 0 max_attention_logit = None - num_microbatches = get_num_microbatches() - seqlen_sum_this_global_batch = None - seqlen_squared_sum_this_global_batch = None else: ft_integration.on_training_step_start() ( @@ -4348,9 +4063,6 @@ def trace_handler(p): grad_norm, num_zeros_in_grad, max_attention_logit, - num_microbatches, - seqlen_sum_this_global_batch, - seqlen_squared_sum_this_global_batch, ) = train_step( forward_step_func, train_data_iterator, @@ -4454,22 +4166,14 @@ def trace_handler(p): else: assert num_skipped_samples_in_batch == 0 args.skipped_train_samples += num_skipped_samples_in_batch - if getattr(config, 'sequence_packing_scheduler', None) is not None and not args.skip_train: - # The scheduler computed these from the real sequence lengths before - # CP padding and rerouting, so use them directly for FLOPs accounting. - assert seqlen_sum_this_global_batch is not None - assert seqlen_squared_sum_this_global_batch is not None - total_real_tokens_in_batch = seqlen_sum_this_global_batch - seqlen_squared_sum_in_batch = seqlen_squared_sum_this_global_batch - else: - # Drain the per-iteration packed-sequence stats so the FLOPs computation - # reflects THD per-chunk causal attention AND excludes padding tokens - # from token-linear work. Returns ``(None, None)`` for unpacked BSHD - # runs (no collective issued), letting ``num_floating_point_operations`` - # fall back to its closed-form defaults. - total_real_tokens_in_batch, seqlen_squared_sum_in_batch = ( - consume_seqlen_stats_in_iteration() - ) + # Drain the per-iteration packed-sequence stats so the FLOPs computation + # reflects THD per-chunk causal attention AND excludes padding tokens + # from token-linear work. Returns ``(None, None)`` for unpacked BSHD + # runs (no collective issued), letting ``num_floating_point_operations`` + # fall back to its closed-form defaults. + total_real_tokens_in_batch, seqlen_squared_sum_in_batch = ( + consume_seqlen_stats_in_iteration() + ) num_floating_point_operations_in_batch = num_floating_point_operations( args, batch_size, @@ -4508,7 +4212,6 @@ def trace_handler(p): is_first_iteration=is_first_iteration, seqlen_squared_sum_in_batch=seqlen_squared_sum_in_batch, total_real_tokens_in_batch=total_real_tokens_in_batch, - num_microbatches=num_microbatches, ) is_first_iteration = False @@ -4612,6 +4315,13 @@ def trace_handler(p): if should_exit: break + # Early-exit paths (exit-duration / exit-interval / signal handler) sys.exit() + # below before the normal-path logging, so record the train-loop finish time here. + if should_exit: + one_logger and one_logger.log_metrics( + {'app_train_loop_finish_time': one_logger_utils.get_timestamp_in_ms()} + ) + # Destroy CUDA Graphs. if args.cuda_graph_impl == "transformer_engine" and cuda_graph_helper.graphs_created(): cuda_graph_helper.delete_cuda_graphs() @@ -4661,6 +4371,9 @@ def trace_handler(p): nccl_allocator.deregister_mem_pool( buf.nccl_mem_pool, buf.data_parallel_group ) + one_logger and one_logger.log_metrics( + {'app_finish_time': one_logger_utils.get_timestamp_in_ms()} + ) wandb_writer = get_wandb_writer() if wandb_writer: wandb_writer.finish() @@ -4705,7 +4418,12 @@ def evaluate( # make validation batch size independent from training batch size eval_batch_size = args.eval_global_batch_size eval_micro_batch_size = args.eval_micro_batch_size - eval_num_microbatches = eval_batch_size // (eval_micro_batch_size * args.data_parallel_size) + # data_parallel_size excludes the GTP-remat axis (it's folded into total_model_size at + # arguments.py:446); each gtp-remat peer consumes a distinct microbatch, so include it in the + # global sample breadth we divide out to recover the microbatch count. + eval_num_microbatches = eval_batch_size // ( + eval_micro_batch_size * args.data_parallel_size * args.gtp_weight_remat_size + ) forward_backward_func = get_forward_backward_func(schedule_pg_collection=pg_collection) # Reductions source per-rank groups from the model (encoder rank -> encoder groups). eval_pgc = get_attr_wrapped_model(model[0], "pg_collection") @@ -4750,24 +4468,11 @@ def evaluate( # Don't care about timing during evaluation config.timers = None ft_integration.on_eval_step_start() - if getattr(config, 'sequence_packing_scheduler', None) is not None: - try: - (packed_data_iterator, scheduled_eval_num_microbatches, _, _) = ( - wrap_data_iterator(data_iterator, config, eval_num_microbatches) - ) - except StopIteration: - # Validation data iterator exhausted, stop evaluation early. - ft_integration.on_eval_step_end() - config.timers = get_timers() - break - else: - packed_data_iterator = data_iterator - scheduled_eval_num_microbatches = eval_num_microbatches loss_dicts = forward_backward_func( forward_step_func=forward_step_func, - data_iterator=packed_data_iterator, + data_iterator=data_iterator, model=model, - num_microbatches=scheduled_eval_num_microbatches, + num_microbatches=eval_num_microbatches, seq_length=args.seq_length, micro_batch_size=eval_micro_batch_size, decoder_seq_length=args.decoder_seq_length, @@ -5255,40 +4960,3 @@ def should_disable_forward_pre_hook(args): ) and args.overlap_param_gather ) - - -def _get_thd_sequence_length_upper_bound(args): - """Return the padded per-sample THD length upper bound used for graph sizing.""" - max_sequence_length = getattr(args, "seq_length", None) - mock_config_spec = None - if getattr(args, "use_varlen_dataset", False): - mock_config_spec = getattr(args, "varlen_mock_dataset_config_json", None) - elif getattr(args, "sft", False): - mock_config_spec = getattr(args, "sft_mock_dataset_config_json", None) - - if mock_config_spec is not None: - from megatron.training.datasets.utils import load_json_arg - - mock_config = load_json_arg(mock_config_spec) - if isinstance(mock_config, dict) and mock_config.get("max_seq_len") is not None: - max_sequence_length = int(mock_config["max_seq_len"]) - - if max_sequence_length is None: - return None - - if getattr(args, "seq_length", None) is not None: - max_sequence_length = min(int(max_sequence_length), int(args.seq_length)) - - cp_size = int(getattr(args, "context_parallel_size", 1) or 1) - if getattr(args, "dynamic_context_parallel", False): - cp_pad = int(getattr(args, "data_parallel_size", 1) or 1) * cp_size * 2 - else: - cp_pad = cp_size * 2 if cp_size > 1 else 1 - - sp_pad = ( - int(getattr(args, "tensor_model_parallel_size", 1) or 1) - if getattr(args, "sequence_parallel", False) - else 1 - ) - pad_granularity = cp_pad * sp_pad - return int(math.ceil(max_sequence_length / pad_granularity) * pad_granularity) diff --git a/megatron/training/utils/__init__.py b/megatron/training/utils/__init__.py index d0a01b6c65d..306c8c32bdc 100644 --- a/megatron/training/utils/__init__.py +++ b/megatron/training/utils/__init__.py @@ -14,6 +14,7 @@ has_nvrx_checkpointing_async_support, has_nvrx_installed, is_first_or_last_pipeline_stage, + is_gtp_remat_active, is_hybrid_model, is_last_rank, is_rank0, diff --git a/megatron/training/utils/common_utils.py b/megatron/training/utils/common_utils.py index dfd77ec7220..c5e69413853 100644 --- a/megatron/training/utils/common_utils.py +++ b/megatron/training/utils/common_utils.py @@ -15,7 +15,7 @@ from megatron.core._rank_utils import safe_get_rank as _safe_get_rank from megatron.core._slurm_utils import resolve_slurm_local_rank from megatron.core.dist_checkpointing.strategies.nvrx import has_nvrx_async_support -from megatron.core.msc_utils import MultiStorageClientFeature, open_file +from megatron.core.msc_utils import maybe_msc try: from transformer_engine.pytorch.optimizers import multi_tensor_applier, multi_tensor_l2norm @@ -49,6 +49,34 @@ from megatron.training import get_adlr_autoresume, get_args, get_timers +def _compute_norm_2(params_list): + """Compute squared L2 norm of a list of tensors. Returns a CUDA scalar.""" + if len(params_list) > 0: + dummy_overflow_buf = torch.tensor([0], dtype=torch.int, device='cuda') + norm, _ = multi_tensor_applier( + multi_tensor_l2norm, dummy_overflow_buf, [params_list], False + ) + return norm * norm + return torch.zeros((1,), dtype=torch.float32, device='cuda') + + +def _get_param_data(param, force_create_fp32_copy, bf16): + """Extract the appropriate data tensor from a param for norm computation. + + Returns (data_tensor, is_sharded) where is_sharded indicates the param has + a sharded main_param from the distributed optimizer. + """ + if bf16: + if not force_create_fp32_copy and hasattr(param, 'main_param'): + if getattr(param, 'main_param_sharded', False): + if param.main_param is not None: + return param.main_param, True + return None, True + return param.main_param, False + return param.data.float(), False + return param.data, False + + def calc_params_l2_norm(model, force_create_fp32_copy=False): """Calculate l2 norm of parameters""" args = get_args() @@ -71,129 +99,118 @@ def calc_params_l2_norm(model, force_create_fp32_copy=False): return calc_dtensor_params_l2_norm(params) - # Seperate moe and dense params - params_data = [] - moe_params_data = [] - sharded_params_data = [] - data_parallel_group = None + # 8 buckets: 4 categories × (non-sharded, sharded optimizer main_param). + # Each category needs different reduction groups. + params_data = [] # Dense, non-sharded + sharded_params_data = [] # Dense, sharded → reduce over dp_cp + gtp_params_data = [] # GTP_remat, non-sharded + gtp_sharded_params_data = [] # GTP_remat, sharded → reduce over dp_cp + moe_params_data = [] # MoE, non-sharded + moe_sharded_params_data = [] # MoE, sharded → reduce over expert_dp + moe_gtp_params_data = [] # MoE-GTP_remat, non-sharded + moe_gtp_sharded_params_data = [] # MoE-GTP_remat sharded → expert_dp + + gtp_rank = mpu.get_gtp_weight_remat_rank() + egtp_rank = mpu.get_expert_gtp_weight_remat_rank() + tp_group = mpu.get_tensor_model_parallel_group() + expert_tp_group = mpu.get_expert_tensor_parallel_group() for model_chunk in model: for param in model_chunk.parameters(): - data_parallel_group = get_data_parallel_group_if_dtensor(param, data_parallel_group) - is_not_tp_duplicate = param_is_not_tensor_parallel_duplicate(param) - if not is_not_tp_duplicate: + is_gtp = getattr(param, 'is_gtp_weight_remat', False) + + # Filter TP duplicates. GTP_remat params are always unique across TP ranks + # so skip this check for them. + if not is_gtp and not param_is_not_tensor_parallel_duplicate( + param, tp_group=tp_group, expert_tp_group=expert_tp_group + ): continue - assert is_not_tp_duplicate - if not getattr(param, 'allreduce', True): + is_expert = not getattr(param, 'allreduce', True) + + # Filter GTP_remat duplicates: non-GTP_remat params replicate across GTP_remat ranks. + if is_expert: + if not is_gtp and egtp_rank != 0: + continue + else: + if not is_gtp and gtp_rank != 0: + continue + + # Route to the correct bucket. + if is_expert: assert param_is_not_shared(param) param = to_local_if_dtensor(param) - if args.bf16: - if not force_create_fp32_copy and hasattr(param, 'main_param'): - if getattr(param, 'main_param_sharded', False): - if param.main_param is not None: - sharded_params_data.append(param.main_param) - else: - moe_params_data.append(param.main_param) - else: - # Fallback to original logic of making a fp32 copy of the - # parameter if `.main_param` attribute is not available. - moe_params_data.append(param.data.float()) + data, is_sharded = _get_param_data(param, force_create_fp32_copy, args.bf16) + if data is None: + continue + if is_gtp: + (moe_gtp_sharded_params_data if is_sharded else moe_gtp_params_data).append( + data + ) else: - moe_params_data.append(param.data) + (moe_sharded_params_data if is_sharded else moe_params_data).append(data) else: if param_is_not_shared(param): param = to_local_if_dtensor(param) - if args.bf16: - if not force_create_fp32_copy and hasattr(param, 'main_param'): - if getattr(param, 'main_param_sharded', False): - if param.main_param is not None: - sharded_params_data.append(param.main_param) - else: - params_data.append(param.main_param) - else: - # Fallback to original logic of making a fp32 copy of the - # parameter if `.main_param` attribute is not available. - params_data.append(param.data.float()) + data, is_sharded = _get_param_data(param, force_create_fp32_copy, args.bf16) + if data is None: + continue + if is_gtp: + (gtp_sharded_params_data if is_sharded else gtp_params_data).append(data) else: - params_data.append(param.data) - - # Calculate norm. - dummy_overflow_buf = torch.tensor([0], dtype=torch.int, device='cuda') - if len(params_data) > 0: - norm, _ = multi_tensor_applier( - multi_tensor_l2norm, dummy_overflow_buf, [params_data], False # no per-parameter norm. - ) - norm_2 = norm * norm - else: - norm_2 = torch.zeros((1,), dtype=torch.float32, device='cuda') - - if data_parallel_group is not None: - torch.distributed.all_reduce( - norm_2, op=torch.distributed.ReduceOp.SUM, group=data_parallel_group - ) - - # Add norm contribution from params with sharded main_params. These norms need to be - # accumulated across the DP group since the main parameters are sharded because - # of distributed optimizer. - if len(sharded_params_data) > 0: - dummy_overflow_buf = torch.tensor([0], dtype=torch.int, device='cuda') - sharded_norm, _ = multi_tensor_applier( - multi_tensor_l2norm, - dummy_overflow_buf, - [sharded_params_data], - False, # no per-parameter norm. - ) - sharded_norm_2 = sharded_norm * sharded_norm - else: - sharded_norm_2 = torch.zeros((1,), dtype=torch.float32, device='cuda') - # Sum over all DP groups, including CP since distributed optimizer state is - # sharded jointly over DP+CP. - torch.distributed.all_reduce( + (sharded_params_data if is_sharded else params_data).append(data) + + # --- Compute local norm^2 for each bucket --- + params_norm_2 = _compute_norm_2(params_data) + sharded_norm_2 = _compute_norm_2(sharded_params_data) + gtp_norm_2 = _compute_norm_2(gtp_params_data) + gtp_sharded_norm_2 = _compute_norm_2(gtp_sharded_params_data) + moe_norm_2 = _compute_norm_2(moe_params_data) + moe_sharded_norm_2 = _compute_norm_2(moe_sharded_params_data) + moe_gtp_norm_2 = _compute_norm_2(moe_gtp_params_data) + moe_gtp_sharded_norm_2 = _compute_norm_2(moe_gtp_sharded_params_data) + + def _sum_reduce(tensor, group): + torch.distributed.all_reduce(tensor, op=torch.distributed.ReduceOp.SUM, group=group) + + # --- Sharded optimizer DP reductions (each category uses its own group) --- + # Reduce over the gtp_remat-EXCLUDED replicate group (with_gtp_remat=False): the model-parallel + # reduce below already spans the gtp_remat axis, so a gtp_remat-inclusive group here would + # over-count by gtp_remat. No-op for non-GTP_remat runs. + _sum_reduce( sharded_norm_2, - op=torch.distributed.ReduceOp.SUM, - group=mpu.get_data_parallel_group(with_context_parallel=True), + mpu.get_data_parallel_group(with_context_parallel=True, with_gtp_remat=False), ) - norm_2 += sharded_norm_2 - - # Add norm contribution from expert layers in MoEs. - if len(moe_params_data) > 0: - moe_norm, _ = multi_tensor_applier( - multi_tensor_l2norm, - dummy_overflow_buf, - [moe_params_data], - False, # no per-parameter norm. - ) - moe_norm_2 = moe_norm * moe_norm + _sum_reduce( + gtp_sharded_norm_2, + mpu.get_data_parallel_group(with_context_parallel=True, with_gtp_remat=False), + ) + _sum_reduce(moe_sharded_norm_2, mpu.get_expert_data_parallel_group(with_gtp_remat=False)) + _sum_reduce(moe_gtp_sharded_norm_2, mpu.get_expert_data_parallel_group(with_gtp_remat=False)) - # Account for MoE norm even if current rank doesn't have any expert params to prevent - # hang in models with un-even numbers of MoE layers. - # See details in https://gitlab-master.nvidia.com/ADLR/megatron-lm/-/issues/409 - else: - moe_norm_2 = torch.zeros_like(norm_2) + # --- Combine dense + GTP_remat norms --- + # model_parallel group = TP×GTP_remat×PP, so GTP_remat reduction is implicit. + norm_2 = params_norm_2 + sharded_norm_2 + gtp_norm_2 + gtp_sharded_norm_2 - # Reduce norm across model parallel groups (dense and expert). - # Dense params should sum across all model-parallel GPUs (tensor + pipeline). + # --- Combine MoE + MoE-GTP_remat norms --- + # expert_model_parallel = TP×EP×PP (does NOT include EGTP_remat), so we need + # an explicit EGTP_remat reduction for MoE-GTP_remat before the model-parallel reduce. + moe_gtp_combined_norm_2 = moe_gtp_norm_2 + moe_gtp_sharded_norm_2 + _sum_reduce(moe_gtp_combined_norm_2, mpu.get_expert_gtp_weight_remat_group()) + moe_total_norm_2 = moe_norm_2 + moe_sharded_norm_2 + moe_gtp_combined_norm_2 + + # --- Model-parallel reductions --- dense_reduce_group = mpu.get_model_parallel_group() - ranks_in_dense_reduce_group = torch.distributed.get_process_group_ranks(dense_reduce_group) - # Expert params should sum across all model-parallel GPUs (expert + tensor + pipeline). expert_reduce_group = mpu.get_expert_tensor_model_pipeline_parallel_group() + ranks_in_dense_reduce_group = torch.distributed.get_process_group_ranks(dense_reduce_group) ranks_in_expert_reduce_group = torch.distributed.get_process_group_ranks(expert_reduce_group) - # If dense and expert reduce groups are the same, sum then reduce. if ranks_in_dense_reduce_group == ranks_in_expert_reduce_group: - norm_2 += moe_norm_2 - torch.distributed.all_reduce( - norm_2, op=torch.distributed.ReduceOp.SUM, group=dense_reduce_group - ) - # If dense and expert reduce groups are different, reduce then sum. + norm_2 += moe_total_norm_2 + _sum_reduce(norm_2, dense_reduce_group) else: - torch.distributed.all_reduce( - norm_2, op=torch.distributed.ReduceOp.SUM, group=dense_reduce_group - ) - torch.distributed.all_reduce( - moe_norm_2, op=torch.distributed.ReduceOp.SUM, group=expert_reduce_group - ) - norm_2 += moe_norm_2 + _sum_reduce(norm_2, dense_reduce_group) + _sum_reduce(moe_total_norm_2, expert_reduce_group) + norm_2 += moe_total_norm_2 return norm_2.item() ** 0.5 @@ -466,6 +483,14 @@ def is_hybrid_model(args): return args.hybrid_layer_pattern is not None +def is_gtp_remat_active(args): + """Returns True if GTP weight-remat is enabled on the decoder or expert axis.""" + return ( + getattr(args, 'gtp_weight_remat_size', 1) > 1 + or getattr(args, 'expert_gtp_weight_remat_size', 1) > 1 + ) + + def is_first_or_last_pipeline_stage(vp_stage): """Return True if on first or last pipeline stage, taking into account virtual pipeline parallelism.""" @@ -498,14 +523,14 @@ def get_blend_and_blend_per_split(args): if use_data_path: if args.data_args_path is not None: assert args.data_path is None - with open_file(args.data_args_path, 'r') as f: + with maybe_msc.open(args.data_args_path, 'r') as f: blend = get_blend_from_list(f.read().split()) else: assert args.data_path is not None blend = get_blend_from_list(args.data_path) elif use_per_split_data_path: if args.per_split_data_args_path is not None: - with open_file(args.per_split_data_args_path, 'r') as f: + with maybe_msc.open(args.per_split_data_args_path, 'r') as f: per_split_data_args = json.load(f) # Each element in blend_per_split should be a list of files (and optional # weights), so split string if needed. diff --git a/pretrain_gpt.py b/pretrain_gpt.py index 81915c0a579..cdd511d0b6d 100644 --- a/pretrain_gpt.py +++ b/pretrain_gpt.py @@ -29,6 +29,7 @@ from megatron.core.datasets.data_schedule import get_batch_on_this_rank_for_sequence_packing from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset from megatron.core.enums import ModelType +from megatron.core.package_info import __version__ as mcore_version from megatron.core.models.gpt import GPTModel from megatron.core.packed_seq_params import ( PackedSeqParams, @@ -42,7 +43,9 @@ from megatron.core.utils import ( StragglerDetector, get_attr_wrapped_model, + get_te_version, get_thd_batch_on_this_cp_rank, + get_torch_version, ) from megatron.training import ( get_args, @@ -52,7 +55,7 @@ print_rank_0, set_startup_timestamps, ) -from megatron.training.argument_utils import pretrain_cfg_container_from_args +from megatron.training.argument_utils import gpt_config_from_args, pretrain_cfg_container_from_args from megatron.training.arguments import core_transformer_config_from_args, parse_and_validate_args from megatron.training.datasets.fim_dataset import GPTFIMDataset, GPTFIMDatasetConfig from megatron.training.datasets.sft_dataset import MockSFTDataset, SFTDataset @@ -68,6 +71,8 @@ try: from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.loss_func import loss_func as loss_func_modelopt + from megatron.post_training.model_builder import ModelOptModelConfig + from megatron.post_training.utils import maybe_enable_modelopt has_nvidia_modelopt = True except ImportError: @@ -544,6 +549,10 @@ def get_embedding_ranks(pp_ranks: List[int]): # Timestamp right after entering __main__ block (after all imports/library setup) _MAIN_ENTRY_TIME = time.time() + print_rank_0(f'> PyTorch version ................ {get_torch_version()}') + print_rank_0(f'> Megatron-Core version .......... {mcore_version}') + print_rank_0(f'> Transformer Engine version ... {get_te_version()}') + # Register startup timestamps for timing report in pretrain() set_startup_timestamps(program_start=_PROGRAM_START_TIME, main_entry=_MAIN_ENTRY_TIME) @@ -557,7 +566,13 @@ def get_embedding_ranks(pp_ranks: List[int]): extra_args_provider=add_modelopt_args if has_nvidia_modelopt else None, args_defaults={'tokenizer_type': 'GPT2BPETokenizer'}, ) - full_config = pretrain_cfg_container_from_args(args) + if has_nvidia_modelopt: + maybe_enable_modelopt(args) + if has_nvidia_modelopt and getattr(args, "modelopt_enabled", False): + model_cfg = gpt_config_from_args(args, model_config_cls=ModelOptModelConfig) + else: + model_cfg = gpt_config_from_args(args) + full_config = pretrain_cfg_container_from_args(args, model_cfg) pretrain( full_config, train_valid_test_datasets_provider, diff --git a/pretrain_hybrid.py b/pretrain_hybrid.py index 5fdc7d8645a..0066796d2bb 100644 --- a/pretrain_hybrid.py +++ b/pretrain_hybrid.py @@ -28,6 +28,7 @@ from megatron.core.datasets.data_schedule import get_batch_on_this_rank_for_sequence_packing from megatron.core.datasets.gpt_dataset import GPTDataset, GPTDatasetConfig, MockGPTDataset from megatron.core.enums import ModelType +from megatron.core.package_info import __version__ as mcore_version from megatron.core.models.hybrid.hybrid_model import HybridModel from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.parallel_state import ( @@ -45,6 +46,8 @@ get_attr_wrapped_model, get_batch_on_this_cp_rank, get_batch_on_this_tp_rank, + get_te_version, + get_torch_version, ) from megatron.training import ( get_args, @@ -68,6 +71,8 @@ try: from megatron.post_training.arguments import add_modelopt_args from megatron.post_training.loss_func import loss_func as loss_func_modelopt + from megatron.post_training.model_builder import ModelOptHybridModelConfig + from megatron.post_training.utils import maybe_enable_modelopt has_nvidia_modelopt = True except ImportError: @@ -445,6 +450,10 @@ def train_valid_test_datasets_provider(train_val_test_num_samples, vp_stage=None # Timestamp right after entering __main__ block (after all imports/library setup) _MAIN_ENTRY_TIME = time.time() + print_rank_0(f'> PyTorch version ................ {get_torch_version()}') + print_rank_0(f'> Megatron-Core version .......... {mcore_version}') + print_rank_0(f'> Transformer Engine version ... {get_te_version()}') + # Register startup timestamps for timing report in pretrain() set_startup_timestamps(program_start=_PROGRAM_START_TIME, main_entry=_MAIN_ENTRY_TIME) @@ -458,7 +467,12 @@ def train_valid_test_datasets_provider(train_val_test_num_samples, vp_stage=None extra_args_provider=add_modelopt_args if has_nvidia_modelopt else None, args_defaults={'tokenizer_type': 'GPT2BPETokenizer'}, ) - model_cfg = hybrid_config_from_args(args) + if has_nvidia_modelopt: + maybe_enable_modelopt(args) + if has_nvidia_modelopt and getattr(args, "modelopt_enabled", False): + model_cfg = hybrid_config_from_args(args, model_config_cls=ModelOptHybridModelConfig) + else: + model_cfg = hybrid_config_from_args(args) full_config = pretrain_cfg_container_from_args(args, model_cfg) pretrain( full_config, diff --git a/skills/mcore-migrate-gpt-to-hybrid/SKILL.md b/skills/mcore-migrate-gpt-to-hybrid/SKILL.md new file mode 100644 index 00000000000..486519ecfb9 --- /dev/null +++ b/skills/mcore-migrate-gpt-to-hybrid/SKILL.md @@ -0,0 +1,43 @@ +--- +name: mcore-migrate-gpt-to-hybrid +description: Migration guide for moving Megatron Core GPTModel checkpoints, model providers, training commands, and layer mappings to HybridModel. +license: Apache-2.0 +when_to_use: Migrating or reviewing a GPTModel checkpoint or training workflow for HybridModel; choosing or reviewing a hybrid layer pattern; running gpt_hybrid_conversion.py; loading a converted checkpoint; diagnosing GPT-to-Hybrid migration issues; 'migrate GPTModel to HybridModel', 'convert GPT checkpoint to HybridModel', 'hybrid layer pattern'. +metadata: + author: Philip Petrakian +--- + +# GPTModel to HybridModel Migration + +## Answer-First Migration Guidance + +- The canonical source is + [`docs/user-guide/hybrid-model-migration.md`](../../docs/user-guide/hybrid-model-migration.md). +- Read the canonical document completely before answering, planning, reviewing, + editing, converting, or training. +- Keep migration behavior, commands, mappings, prerequisites, limitations, and + validation in the canonical document only. Do not duplicate them in this + skill. + +--- + +## Workflow + +1. Pull the task artifact first: checkpoint metadata, model provider or config, + training command, conversion log, diff, or failure output. +2. Read the canonical migration document completely. +3. Follow only the relevant document sections. Do not invent an unsupported + migration path or silently change the target architecture. +4. Validate the result proportionately, invoking the relevant repository build + and testing skills when applicable. +5. Report the outcome and link the canonical document for human readers. + +--- + +## Documentation Drift + +If the implementation and migration guide disagree: + +1. Report the discrepancy before continuing. +2. If the task authorizes a correction, update the canonical document first. +3. Do not add a competing migration rule to this skill. diff --git a/skills/nightly-sync/SKILL.md b/skills/nightly-sync/SKILL.md index d350d4a7a6f..cd3b85f2e0a 100644 --- a/skills/nightly-sync/SKILL.md +++ b/skills/nightly-sync/SKILL.md @@ -232,13 +232,19 @@ Run on ALL changed Python files (relative to `origin/dev`), in this order: 4. `pylint` on changed `megatron/core/` files — fix missing-docstring and line-too-long violations before pushing -### Pre-push invariant checks +### Pre-push advisory checks Before every `git push` in this workflow (the initial push in Phase 1 -AND every fix-push in Phase 3), run these bash checks. If any fails, -fix the condition and re-check before pushing: +AND every fix-push in Phase 3), run these bash checks as guidance. They +must never block the push. Review every finding: fix genuine merge accidents +and document intentional main removals or formatting/reordering false positives +in the PR body. ```bash +set +e +( +set -euo pipefail + MERGE_COMMIT=$(git rev-list --min-parents=2 --max-count=1 HEAD || true) if [ -n "$MERGE_COMMIT" ]; then DEV_REF="${MERGE_COMMIT}^1" @@ -250,9 +256,8 @@ fi # 1. CODEOWNERS must be identical to dev's. if ! git diff --quiet "$DEV_REF" HEAD -- .github/CODEOWNERS; then - echo "ABORT: .github/CODEOWNERS differs from dev. Restore with:" + echo "WARNING: .github/CODEOWNERS differs from dev. Restore with:" echo " git checkout $DEV_REF -- .github/CODEOWNERS" - exit 1 fi # 2. Dependency-management triple must be identical to dev's. @@ -260,7 +265,7 @@ for f in pyproject.toml uv.lock docker/Dockerfile.ci.dev; do if ! git diff --quiet "$DEV_REF" HEAD -- "$f"; then # pyproject.toml is allowed to differ ONLY for git source reconciliation # (new [tool.uv.sources] entries from main). If you intentionally edited - # it for that reason, bypass this check by re-running with $f skipped. + # it for that reason, document the reconciliation in the PR body. echo "WARNING: $f differs from dev" fi done @@ -291,7 +296,7 @@ done INTENTIONAL_OVERRIDE_REGEX='^(megatron/training/training\.py|megatron/training/initialize\.py|megatron/training/utils\.py|megatron/training/datasets/data_samplers\.py|megatron/core/optimizer/layer_wise_optimizer\.py)$' SKIP_REGEX='^(pyproject\.toml|uv\.lock|docker/Dockerfile\.ci\.dev|\.github/CODEOWNERS)$' -VIOLATIONS=0 +FINDINGS=0 for f in $(git diff --name-only "$DEV_REF"..HEAD \ -- '*.py' '*.md' '*.yaml' '*.yml' '*.toml' \ '*.sh' '*.cpp' '*.cu' '*.h' \ @@ -310,25 +315,36 @@ for f in $(git diff --name-only "$DEV_REF"..HEAD \ if [ -n "$missing" ]; then echo "=== $f ===" printf '%s\n' "$missing" - VIOLATIONS=$((VIOLATIONS + $(printf '%s\n' "$missing" | grep -c .))) + FINDINGS=$((FINDINGS + $(printf '%s\n' "$missing" | grep -c .))) fi done -if [ "$VIOLATIONS" -gt 0 ]; then - echo "ABORT: $VIOLATIONS dev-only line(s) dropped by the merge. For each:" +if [ "$FINDINGS" -gt 0 ]; then + echo "WARNING: $FINDINGS potential dev-only line removal(s) detected. For each:" echo " (a) MAIN INTENTIONALLY REMOVED — find the specific commit in" echo " 'git log origin/main -- ' that removed it; document the" echo " SHA in the PR body, then the drop is acceptable." echo " (b) MERGE ACCIDENT — main never explicitly touched that line." echo " RESTORE the dev line (Edit/Write to put it back)." echo "Default to (b); only declare (a) with a specific main commit as evidence." - exit 1 + echo "This audit is advisory; continue the push after reviewing the findings." +fi + +echo "nightly-sync pre-push guidance complete" +) +GUIDANCE_STATUS=$? +if [ "$GUIDANCE_STATUS" -ne 0 ]; then + echo "WARNING: nightly-sync pre-push guidance failed with status $GUIDANCE_STATUS; allowing the push to continue." fi +exit 0 ``` -The CODEOWNERS check and the dev-feature preservation audit are HARD -aborts — never push if either fails. The dep-triple check is a warning -because git-source reconciliation can produce legitimate diffs there. +All pre-push findings are advisory. The hook must return success even when it +finds a CODEOWNERS difference, potential dev-feature removal, dependency-triple +difference, or an internal audit error. The underlying policies still apply: +restore accidental changes, preserve dev-only features, and document exact main +commits for intentional removals. A warning by itself is never a reason to stop +the workflow or request authorization to continue. Recent regressions the dev-feature audit would have flagged (all "merge accident" type from #4659 and #4716): @@ -393,7 +409,12 @@ Phase 3 step 4 and the two-commit policy in Rules). in the PR body so reviewers see it at a glance. 3. List of files where main's version was taken over the merge 4. List of files that were deleted in dev but restored (and why) - 5. The remerge-diff output (`git show --remerge-diff HEAD` on the merge + 5. Disposition of every pre-push advisory finding, including CODEOWNERS or + dependency-triple differences and potential dev-feature removals. Record + whether each was corrected, intentional (with the exact commit or reason), + or a formatting/reordering false positive. State explicitly if there were + no findings. + 6. The remerge-diff output (`git show --remerge-diff HEAD` on the merge commit) so reviewers can inspect ONLY the conflict resolutions. If the output is very long, summarize conflicts by file and put the full diff in a collapsed `
` block. If git is too old for `--remerge-diff`, diff --git a/tests/functional_tests/dynamic_inference_functional_tests.md b/tests/functional_tests/dynamic_inference_functional_tests.md index 2cc84670484..b1453700cb0 100644 --- a/tests/functional_tests/dynamic_inference_functional_tests.md +++ b/tests/functional_tests/dynamic_inference_functional_tests.md @@ -53,7 +53,7 @@ CLI flags below are verified to exist in `megatron/training/arguments.py` and/or |---|---|---|---| | Enable prefix caching | `--inference-dynamic-batching-prefix-caching` | off | Reuse KV blocks for shared prompt prefixes | | Eviction policy | `--inference-dynamic-batching-prefix-caching-eviction-policy {ref_zero, lru}` | `ref_zero` | Block reclamation strategy | -| Coordinator routing | `--inference-dynamic-batching-prefix-caching-coordinator-policy {longest_prefix, first_prefix_block, round_robin}` | `first_prefix_block` | Multi-rank request routing | +| Coordinator routing | `--inference-dynamic-batching-prefix-caching-coordinator-policy {longest_prefix, first_prefix_block, load_balanced}` | `load_balanced` | Multi-rank request routing | | Routing alpha | `--inference-dynamic-batching-prefix-caching-routing-alpha` | 0.5 | 0=load-balance, 1=prefix-affinity | | Mamba state cache | `--inference-dynamic-batching-prefix-caching-mamba-gb` | — | GPU memory for Mamba hybrid block states | @@ -306,7 +306,7 @@ User selected the **recommended cut** of 6 tests + drift fix + cw-dfw golden gen **Decision (2026-05-12):** User selected Tier 2 (18 tests). Substitutions vs. original Tier 2 pitch: - CP-parallelism tests **dropped** — dynamic inference doesn't support CP (issue #6). -- Added 2 ZMQ-coordinator tests (longest_prefix + round_robin policies). +- Added 2 ZMQ-coordinator tests (longest_prefix + load_balanced policies). | # | Test | model_config | recipe | run | golden | pytest verified | |---|---|---|---|---|---|---| @@ -327,7 +327,7 @@ User selected the **recommended cut** of 6 tests + drift fix + cw-dfw golden gen | 15 | `gpt_dynamic_inference_tp4_pp1_ep4_16B_prefix_caching` (MoE) | ✅ | ✅ | ✅ | ✅ committed | ✅ **PASSED** | | 16 | `gpt_dynamic_inference_tp4_pp1_ep4_16B_chunked_prefill` (MoE) | ✅ | ✅ | ✅ | ✅ committed | ⏸️ (not sampled) | | 17 | `gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_longest_prefix_zmq` | ✅ | ✅ | ✅ | ✅ committed | ✅ **PASSED** | -| 18 | `gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_round_robin_zmq` | ✅ | ✅ | ✅ | ✅ committed | ⏸️ (not sampled) | +| 18 | `gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_load_balanced_zmq` | ✅ | ✅ | ✅ | ✅ committed | ⏸️ (not sampled) | **Verification methodology**: ran 4 representative tests (one per category: parallelism, 3-way combo, MoE, DP+ZMQ) without `RECORD_CHECKPOINTS=true` so the pytest comparison actually executes. All 4 reported `test_inference_pipeline PASSED` against their committed goldens. diff --git a/tests/functional_tests/python_test_utils/get_test_results_from_tensorboard_logs.py b/tests/functional_tests/python_test_utils/get_test_results_from_tensorboard_logs.py index 091623b9b84..fcee30e61e6 100644 --- a/tests/functional_tests/python_test_utils/get_test_results_from_tensorboard_logs.py +++ b/tests/functional_tests/python_test_utils/get_test_results_from_tensorboard_logs.py @@ -50,6 +50,7 @@ def collect_train_test_metrics( "lm loss", "num-zeros", "mtp_1 loss", + "total loss", ] } diff --git a/tests/functional_tests/python_test_utils/test_inference_regular_pipeline.py b/tests/functional_tests/python_test_utils/test_inference_regular_pipeline.py index bec93f10675..0e6e0974683 100644 --- a/tests/functional_tests/python_test_utils/test_inference_regular_pipeline.py +++ b/tests/functional_tests/python_test_utils/test_inference_regular_pipeline.py @@ -148,7 +148,8 @@ def test_inference_pipeline( # TODO: Compare liftime_prefill_token_count to groundtruth pass - for request_id, groundtruth_results in output_groundtruth.items(): + for request_id in groundtruth_request_ids: + groundtruth_results = output_groundtruth[request_id] current_results = output_current[request_id] at_least_one_test_loop = False diff --git a/tests/functional_tests/python_test_utils/test_pretraining_regular_pipeline.py b/tests/functional_tests/python_test_utils/test_pretraining_regular_pipeline.py index 68aa0db5622..a5ad326f49d 100644 --- a/tests/functional_tests/python_test_utils/test_pretraining_regular_pipeline.py +++ b/tests/functional_tests/python_test_utils/test_pretraining_regular_pipeline.py @@ -20,6 +20,7 @@ "num-zeros": [common.DeterministicTest(), common.ApproximateTest(atol=0, rtol=0.20)], "generated_tokens": [common.DeterministicTest(), common.ApproximateTest(atol=0, rtol=0.05)], "logprobs": [common.DeterministicTest(), common.ApproximateTest(atol=0, rtol=0.05)], + "total loss": [common.DeterministicTest(), common.ApproximateTest(atol=0, rtol=0.05)], } diff --git a/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp2/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp2/golden_values_dev_dgx_h100.json index 1d75976567a..cb263d48506 100644 --- a/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp2/golden_values_dev_dgx_h100.json +++ b/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp2/golden_values_dev_dgx_h100.json @@ -26,34 +26,34 @@ "20": 10.47714, "21": 10.45276, "22": 10.39141, - "23": "nan", - "24": "nan", - "25": "nan", - "26": "nan", - "27": "nan", - "28": "nan", - "29": "nan", - "30": "nan", - "31": "nan", - "32": "nan", - "33": "nan", - "34": "nan", - "35": "nan", - "36": "nan", - "37": "nan", - "38": "nan", - "39": "nan", - "40": "nan", - "41": "nan", - "42": "nan", - "43": "nan", - "44": "nan", - "45": "nan", - "46": "nan", - "47": "nan", - "48": "nan", - "49": "nan", - "50": "nan" + "23": 10.3972, + "24": 10.35475, + "25": 10.35246, + "26": 10.35041, + "27": 10.31147, + "28": 10.32877, + "29": 10.30861, + "30": 10.14449, + "31": 10.10687, + "32": 10.07856, + "33": 10.10424, + "34": 10.03458, + "35": 10.02764, + "36": 10.01018, + "37": 10.00276, + "38": 9.96156, + "39": 9.89009, + "40": 9.85306, + "41": 9.78426, + "42": 9.71982, + "43": 9.68658, + "44": 9.65517, + "45": 9.6502, + "46": 9.57537, + "47": 9.59269, + "48": 9.58288, + "49": 9.52859, + "50": 9.49558 } }, "num-zeros": { @@ -83,34 +83,34 @@ "20": 2266.0, "21": 2428.0, "22": 2319.0, - "23": "nan", - "24": "nan", - "25": "nan", - "26": "nan", - "27": "nan", - "28": "nan", - "29": "nan", - "30": "nan", - "31": "nan", - "32": "nan", - "33": "nan", - "34": "nan", - "35": "nan", - "36": "nan", - "37": "nan", - "38": "nan", - "39": "nan", - "40": "nan", - "41": "nan", - "42": "nan", - "43": "nan", - "44": "nan", - "45": "nan", - "46": "nan", - "47": "nan", - "48": "nan", - "49": "nan", - "50": "nan" + "23": 2420.0, + "24": 2343.0, + "25": 2235.0, + "26": 2722.0, + "27": 2402.0, + "28": 2568.0, + "29": 1915.0, + "30": 2132.0, + "31": 2699.0, + "32": 2340.0, + "33": 2667.0, + "34": 2775.0, + "35": 2753.0, + "36": 1838.0, + "37": 2756.0, + "38": 2479.0, + "39": 2316.0, + "40": 3061.0, + "41": 3153.0, + "42": 3112.0, + "43": 2808.0, + "44": 3013.0, + "45": 3282.0, + "46": 3037.0, + "47": 3164.0, + "48": 3314.0, + "49": 2706.0, + "50": 2787.0 } }, "mem-allocated-bytes": { @@ -140,34 +140,34 @@ "20": 3433473024.0, "21": 3433473024.0, "22": 3433473024.0, - "23": "nan", - "24": "nan", - "25": "nan", - "26": "nan", - "27": "nan", - "28": "nan", - "29": "nan", - "30": "nan", - "31": "nan", - "32": "nan", - "33": "nan", - "34": "nan", - "35": "nan", - "36": "nan", - "37": "nan", - "38": "nan", - "39": "nan", - "40": "nan", - "41": "nan", - "42": "nan", - "43": "nan", - "44": "nan", - "45": "nan", - "46": "nan", - "47": "nan", - "48": "nan", - "49": "nan", - "50": "nan" + "23": 3433473024.0, + "24": 3433473024.0, + "25": 3433473024.0, + "26": 3433473024.0, + "27": 3433473024.0, + "28": 3433473024.0, + "29": 3433473024.0, + "30": 3433473024.0, + "31": 3433473024.0, + "32": 3433473024.0, + "33": 3433473024.0, + "34": 3433473024.0, + "35": 3433473024.0, + "36": 3433473024.0, + "37": 3433473024.0, + "38": 3433473024.0, + "39": 3433473024.0, + "40": 3433473024.0, + "41": 3433473024.0, + "42": 3433473024.0, + "43": 3433473024.0, + "44": 3433473024.0, + "45": 3433473024.0, + "46": 3433473024.0, + "47": 3433473024.0, + "48": 3433473024.0, + "49": 3433473024.0, + "50": 3433473024.0 } }, "mem-max-allocated-bytes": { @@ -197,34 +197,34 @@ "20": 5707393024.0, "21": 5707393024.0, "22": 5707393024.0, - "23": "nan", - "24": "nan", - "25": "nan", - "26": "nan", - "27": "nan", - "28": "nan", - "29": "nan", - "30": "nan", - "31": "nan", - "32": "nan", - "33": "nan", - "34": "nan", - "35": "nan", - "36": "nan", - "37": "nan", - "38": "nan", - "39": "nan", - "40": "nan", - "41": "nan", - "42": "nan", - "43": "nan", - "44": "nan", - "45": "nan", - "46": "nan", - "47": "nan", - "48": "nan", - "49": "nan", - "50": "nan" + "23": 5707393024.0, + "24": 5707393024.0, + "25": 5707393024.0, + "26": 5707393024.0, + "27": 5707393024.0, + "28": 5707393024.0, + "29": 5707393024.0, + "30": 5707393024.0, + "31": 5707393024.0, + "32": 5707393024.0, + "33": 5707393024.0, + "34": 5707393024.0, + "35": 5707393024.0, + "36": 5707393024.0, + "37": 5707393024.0, + "38": 5707393024.0, + "39": 5707393024.0, + "40": 5707393024.0, + "41": 5707393024.0, + "42": 5707393024.0, + "43": 5707393024.0, + "44": 5707393024.0, + "45": 5707393024.0, + "46": 5707393024.0, + "47": 5707393024.0, + "48": 5707393024.0, + "49": 5707393024.0, + "50": 5707393024.0 } }, "iteration-time": { @@ -232,56 +232,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": "nan", - "2": 5.74712, - "3": 0.49778, - "4": 0.45748, - "5": 0.45659, - "6": 0.45981, - "7": 0.46548, - "8": 0.46542, - "9": 0.46526, - "10": 0.46413, - "11": 0.45692, - "12": 0.46222, - "13": 0.46736, - "14": 0.46657, - "15": 0.46742, - "16": 0.46727, - "17": 0.4733, - "18": 0.469, - "19": 0.45727, - "20": 0.47259, - "21": 0.46632, - "22": 0.46891, - "23": "nan", - "24": "nan", - "25": "nan", - "26": "nan", - "27": "nan", - "28": "nan", - "29": "nan", - "30": "nan", - "31": "nan", - "32": "nan", - "33": "nan", - "34": "nan", - "35": "nan", - "36": "nan", - "37": "nan", - "38": "nan", - "39": "nan", - "40": "nan", - "41": "nan", - "42": "nan", - "43": "nan", - "44": "nan", - "45": "nan", - "46": "nan", - "47": "nan", - "48": "nan", - "49": "nan", - "50": "nan" + "1": 0.0, + "2": 6.38114, + "3": 0.48561, + "4": 0.89062, + "5": 0.43795, + "6": 0.43936, + "7": 1.41561, + "8": 0.87474, + "9": 0.87862, + "10": 0.91586, + "11": 0.44504, + "12": 0.44695, + "13": 0.45378, + "14": 0.45293, + "15": 0.4528, + "16": 0.43931, + "17": 0.43805, + "18": 0.93878, + "19": 0.44812, + "20": 0.44479, + "21": 0.91621, + "22": 0.44682, + "23": 0.44976, + "24": 0.44655, + "25": 0.43297, + "26": 0.43824, + "27": 0.44048, + "28": 1.4425, + "29": 0.44398, + "30": 0.44292, + "31": 0.45131, + "32": 0.45161, + "33": 0.45562, + "34": 0.45364, + "35": 0.43687, + "36": 0.43542, + "37": 0.43901, + "38": 0.4423, + "39": 0.44083, + "40": 0.44685, + "41": 0.46676, + "42": 0.49012, + "43": 0.45932, + "44": 0.46833, + "45": 0.43681, + "46": 0.48333, + "47": 0.43935, + "48": 0.43647, + "49": 0.4394, + "50": 0.48458 } } -} \ No newline at end of file +} diff --git a/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp2/model_config.yaml b/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp2/model_config.yaml index 0dc97066835..5f410958836 100644 --- a/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp2/model_config.yaml +++ b/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp2/model_config.yaml @@ -42,3 +42,4 @@ MODEL_ARGS: --ckpt-format: torch --attention-backend: unfused TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp4_vp2/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp4_vp2/golden_values_dev_dgx_h100.json index 651f9de50bb..909aa0ef055 100644 --- a/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp4_vp2/golden_values_dev_dgx_h100.json +++ b/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp4_vp2/golden_values_dev_dgx_h100.json @@ -17,43 +17,43 @@ "11": 10.49164, "12": 10.47821, "13": 10.47598, - "14": "nan", - "15": "nan", - "16": "nan", - "17": "nan", - "18": "nan", - "19": "nan", - "20": "nan", - "21": "nan", - "22": "nan", - "23": "nan", - "24": "nan", - "25": "nan", - "26": "nan", - "27": "nan", - "28": "nan", - "29": "nan", - "30": "nan", - "31": "nan", - "32": "nan", - "33": "nan", - "34": "nan", - "35": "nan", - "36": "nan", - "37": "nan", - "38": "nan", - "39": "nan", - "40": "nan", - "41": "nan", - "42": "nan", - "43": "nan", - "44": "nan", - "45": "nan", - "46": "nan", - "47": "nan", - "48": "nan", - "49": "nan", - "50": "nan" + "14": 10.48189, + "15": 10.48217, + "16": 10.46292, + "17": 10.45829, + "18": 10.45957, + "19": 10.43356, + "20": 10.45335, + "21": 10.42765, + "22": 10.37288, + "23": 10.3837, + "24": 10.34039, + "25": 10.30874, + "26": 10.32676, + "27": 10.3348, + "28": 10.31238, + "29": 10.20958, + "30": 10.10211, + "31": 10.07247, + "32": 10.04225, + "33": 10.04856, + "34": 9.96979, + "35": 9.96036, + "36": 9.94987, + "37": 9.93538, + "38": 9.91494, + "39": 9.81544, + "40": 9.7735, + "41": 9.73656, + "42": 9.68286, + "43": 9.66796, + "44": 9.64166, + "45": 9.64023, + "46": 9.56948, + "47": 9.60362, + "48": 9.59334, + "49": 9.54843, + "50": 9.50472 } }, "num-zeros": { @@ -74,43 +74,43 @@ "11": 2246.0, "12": 1932.0, "13": 2162.0, - "14": "nan", - "15": "nan", - "16": "nan", - "17": "nan", - "18": "nan", - "19": "nan", - "20": "nan", - "21": "nan", - "22": "nan", - "23": "nan", - "24": "nan", - "25": "nan", - "26": "nan", - "27": "nan", - "28": "nan", - "29": "nan", - "30": "nan", - "31": "nan", - "32": "nan", - "33": "nan", - "34": "nan", - "35": "nan", - "36": "nan", - "37": "nan", - "38": "nan", - "39": "nan", - "40": "nan", - "41": "nan", - "42": "nan", - "43": "nan", - "44": "nan", - "45": "nan", - "46": "nan", - "47": "nan", - "48": "nan", - "49": "nan", - "50": "nan" + "14": 2390.0, + "15": 2034.0, + "16": 2039.0, + "17": 2152.0, + "18": 2153.0, + "19": 2104.0, + "20": 2360.0, + "21": 2119.0, + "22": 2155.0, + "23": 2343.0, + "24": 2293.0, + "25": 2243.0, + "26": 2619.0, + "27": 2396.0, + "28": 2465.0, + "29": 1927.0, + "30": 2561.0, + "31": 2128.0, + "32": 2812.0, + "33": 2801.0, + "34": 2333.0, + "35": 2276.0, + "36": 1626.0, + "37": 2759.0, + "38": 2349.0, + "39": 2650.0, + "40": 1821.0, + "41": 1530.0, + "42": 1737.0, + "43": 1415.0, + "44": 2609.0, + "45": 3604.0, + "46": 3307.0, + "47": 2903.0, + "48": 1855.0, + "49": 3165.0, + "50": 3023.0 } }, "mem-allocated-bytes": { @@ -131,43 +131,43 @@ "11": 2090120192.0, "12": 2090120192.0, "13": 2090120192.0, - "14": "nan", - "15": "nan", - "16": "nan", - "17": "nan", - "18": "nan", - "19": "nan", - "20": "nan", - "21": "nan", - "22": "nan", - "23": "nan", - "24": "nan", - "25": "nan", - "26": "nan", - "27": "nan", - "28": "nan", - "29": "nan", - "30": "nan", - "31": "nan", - "32": "nan", - "33": "nan", - "34": "nan", - "35": "nan", - "36": "nan", - "37": "nan", - "38": "nan", - "39": "nan", - "40": "nan", - "41": "nan", - "42": "nan", - "43": "nan", - "44": "nan", - "45": "nan", - "46": "nan", - "47": "nan", - "48": "nan", - "49": "nan", - "50": "nan" + "14": 2090120192.0, + "15": 2090120192.0, + "16": 2090120192.0, + "17": 2090120192.0, + "18": 2090120192.0, + "19": 2090120192.0, + "20": 2090120192.0, + "21": 2090120192.0, + "22": 2090120192.0, + "23": 2090120192.0, + "24": 2090120192.0, + "25": 2090120192.0, + "26": 2090120192.0, + "27": 2090120192.0, + "28": 2090120192.0, + "29": 2090120192.0, + "30": 2090120192.0, + "31": 2090120192.0, + "32": 2090120192.0, + "33": 2090120192.0, + "34": 2090120192.0, + "35": 2090120192.0, + "36": 2090120192.0, + "37": 2090120192.0, + "38": 2090120192.0, + "39": 2090120192.0, + "40": 2090120192.0, + "41": 2090120192.0, + "42": 2090120192.0, + "43": 2090120192.0, + "44": 2090120192.0, + "45": 2090120192.0, + "46": 2090120192.0, + "47": 2090120192.0, + "48": 2090120192.0, + "49": 2090120192.0, + "50": 2090120192.0 } }, "mem-max-allocated-bytes": { @@ -175,56 +175,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 4420291584.0, + "1": 4420548096.0, "2": 5293134848.0, "3": 5293134848.0, "4": 5293134848.0, "5": 5293134848.0, "6": 5293134848.0, "7": 5293134848.0, - "8": 5293134848.0, - "9": 5293917696.0, - "10": 5293917696.0, - "11": 5293917696.0, - "12": 5293917696.0, - "13": 5293917696.0, - "14": "nan", - "15": "nan", - "16": "nan", - "17": "nan", - "18": "nan", - "19": "nan", - "20": "nan", - "21": "nan", - "22": "nan", - "23": "nan", - "24": "nan", - "25": "nan", - "26": "nan", - "27": "nan", - "28": "nan", - "29": "nan", - "30": "nan", - "31": "nan", - "32": "nan", - "33": "nan", - "34": "nan", - "35": "nan", - "36": "nan", - "37": "nan", - "38": "nan", - "39": "nan", - "40": "nan", - "41": "nan", - "42": "nan", - "43": "nan", - "44": "nan", - "45": "nan", - "46": "nan", - "47": "nan", - "48": "nan", - "49": "nan", - "50": "nan" + "8": 5293919744.0, + "9": 5293919744.0, + "10": 5293919744.0, + "11": 5293919744.0, + "12": 5293919744.0, + "13": 5293919744.0, + "14": 5293919744.0, + "15": 5293919744.0, + "16": 5293919744.0, + "17": 5293919744.0, + "18": 5293919744.0, + "19": 5293919744.0, + "20": 5293919744.0, + "21": 5293919744.0, + "22": 5293919744.0, + "23": 5293919744.0, + "24": 5293919744.0, + "25": 5293919744.0, + "26": 5293919744.0, + "27": 5293919744.0, + "28": 5293919744.0, + "29": 5293919744.0, + "30": 5293919744.0, + "31": 5293919744.0, + "32": 5293919744.0, + "33": 5293919744.0, + "34": 5293919744.0, + "35": 5293919744.0, + "36": 5293919744.0, + "37": 5293919744.0, + "38": 5293919744.0, + "39": 5293919744.0, + "40": 5293919744.0, + "41": 5293919744.0, + "42": 5293919744.0, + "43": 5293919744.0, + "44": 5293919744.0, + "45": 5293919744.0, + "46": 5293919744.0, + "47": 5293919744.0, + "48": 5293919744.0, + "49": 5293919744.0, + "50": 5293919744.0 } }, "iteration-time": { @@ -232,56 +232,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": "nan", - "2": 8.61722, - "3": 0.57542, - "4": 0.56623, - "5": 0.5541, - "6": 0.56688, - "7": 0.54082, - "8": 0.54548, - "9": 0.55074, - "10": 0.55601, - "11": 0.55416, - "12": 0.5472, - "13": 1.6631, - "14": "nan", - "15": "nan", - "16": "nan", - "17": "nan", - "18": "nan", - "19": "nan", - "20": "nan", - "21": "nan", - "22": "nan", - "23": "nan", - "24": "nan", - "25": "nan", - "26": "nan", - "27": "nan", - "28": "nan", - "29": "nan", - "30": "nan", - "31": "nan", - "32": "nan", - "33": "nan", - "34": "nan", - "35": "nan", - "36": "nan", - "37": "nan", - "38": "nan", - "39": "nan", - "40": "nan", - "41": "nan", - "42": "nan", - "43": "nan", - "44": "nan", - "45": "nan", - "46": "nan", - "47": "nan", - "48": "nan", - "49": "nan", - "50": "nan" + "1": 0.0, + "2": 9.61344, + "3": 1.04625, + "4": 2.21477, + "5": 2.14621, + "6": 2.08582, + "7": 1.14391, + "8": 1.65552, + "9": 1.48661, + "10": 1.9757, + "11": 0.54115, + "12": 1.89313, + "13": 0.5415, + "14": 0.54171, + "15": 0.57114, + "16": 0.55481, + "17": 0.52196, + "18": 0.52373, + "19": 0.54214, + "20": 0.52465, + "21": 0.52922, + "22": 0.54449, + "23": 0.52283, + "24": 0.54444, + "25": 0.55507, + "26": 0.53438, + "27": 0.54903, + "28": 1.81392, + "29": 0.52179, + "30": 0.53052, + "31": 0.53473, + "32": 0.55591, + "33": 0.55978, + "34": 0.52699, + "35": 0.53671, + "36": 1.04034, + "37": 0.56419, + "38": 0.56171, + "39": 0.53162, + "40": 0.54177, + "41": 0.54959, + "42": 0.54247, + "43": 0.54086, + "44": 0.54122, + "45": 0.54784, + "46": 0.55077, + "47": 0.54358, + "48": 0.55001, + "49": 0.53455, + "50": 0.52111 } } -} \ No newline at end of file +} diff --git a/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp4_vp2/model_config.yaml b/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp4_vp2/model_config.yaml index 05117c8f4a0..a406edabde6 100644 --- a/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp4_vp2/model_config.yaml +++ b/tests/functional_tests/test_cases/bert/bert_mcore_tp1_pp4_vp2/model_config.yaml @@ -43,3 +43,4 @@ MODEL_ARGS: --ckpt-format: torch --attention-backend: unfused TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2/model_config.yaml b/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2/model_config.yaml index 01ba8adeccf..b457d9b74cc 100644 --- a/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2/model_config.yaml +++ b/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2/model_config.yaml @@ -42,3 +42,4 @@ MODEL_ARGS: --ckpt-format: torch --attention-backend: unfused TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_frozen_resume_torch_dist/model_config.yaml b/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_frozen_resume_torch_dist/model_config.yaml index 680c3c69ea7..df1be1da58a 100644 --- a/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_frozen_resume_torch_dist/model_config.yaml +++ b/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_frozen_resume_torch_dist/model_config.yaml @@ -45,3 +45,4 @@ MODEL_ARGS: --dist-ckpt-strictness: log_all # backward compatibility for TE changes --attention-backend: unfused TEST_TYPE: frozen-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_local_spec/model_config.yaml b/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_local_spec/model_config.yaml index c372de7180a..6316bcff0a6 100644 --- a/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_local_spec/model_config.yaml +++ b/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_local_spec/model_config.yaml @@ -43,3 +43,4 @@ MODEL_ARGS: --ckpt-format: torch --attention-backend: local TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_resume_torch_dist/model_config.yaml b/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_resume_torch_dist/model_config.yaml index 4afcb0c9d47..3099db79789 100644 --- a/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_resume_torch_dist/model_config.yaml +++ b/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_resume_torch_dist/model_config.yaml @@ -45,3 +45,4 @@ MODEL_ARGS: --dist-ckpt-strictness: log_all # backward compatibility for TE changes --attention-backend: unfused TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_resume_torch_dist_local_spec/model_config.yaml b/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_resume_torch_dist_local_spec/model_config.yaml index 8a776e6bfe5..608274462b3 100644 --- a/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_resume_torch_dist_local_spec/model_config.yaml +++ b/tests/functional_tests/test_cases/bert/bert_mcore_tp2_pp2_resume_torch_dist_local_spec/model_config.yaml @@ -46,3 +46,4 @@ MODEL_ARGS: --ckpt-format: torch --attention-backend: local TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/bert/bert_mcore_tp4_pp1/model_config.yaml b/tests/functional_tests/test_cases/bert/bert_mcore_tp4_pp1/model_config.yaml index 15ec6afdebe..10f2b2e7047 100644 --- a/tests/functional_tests/test_cases/bert/bert_mcore_tp4_pp1/model_config.yaml +++ b/tests/functional_tests/test_cases/bert/bert_mcore_tp4_pp1/model_config.yaml @@ -42,3 +42,4 @@ MODEL_ARGS: --ckpt-format: torch --attention-backend: unfused TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/common/ckpt_converter/model_config.yaml b/tests/functional_tests/test_cases/common/ckpt_converter/model_config.yaml index 2ac5db11472..fbafc91b4f8 100644 --- a/tests/functional_tests/test_cases/common/ckpt_converter/model_config.yaml +++ b/tests/functional_tests/test_cases/common/ckpt_converter/model_config.yaml @@ -5,3 +5,4 @@ ENV_VARS: CUBLAS_WORKSPACE_CONFIG: :4096:8 MODEL_ARGS: TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt-nemo/bert-nemo_340m_mr_mbs2_gbs32_mcore_te_tp2_pp2_1N8G/model_config.yaml b/tests/functional_tests/test_cases/gpt-nemo/bert-nemo_340m_mr_mbs2_gbs32_mcore_te_tp2_pp2_1N8G/model_config.yaml index b7edb433b46..612bc41ee17 100644 --- a/tests/functional_tests/test_cases/gpt-nemo/bert-nemo_340m_mr_mbs2_gbs32_mcore_te_tp2_pp2_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt-nemo/bert-nemo_340m_mr_mbs2_gbs32_mcore_te_tp2_pp2_1N8G/model_config.yaml @@ -15,3 +15,4 @@ MODEL_ARGS: data.seq_length: 512 log.log_dir: ${CHECKPOINT_SAVE_PATH} TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt-nemo/gemma2-nemo_2b_mr_mbs1_gbs8_mcore_te_tp4_pp1_cp1_1N8G/model_config.yaml b/tests/functional_tests/test_cases/gpt-nemo/gemma2-nemo_2b_mr_mbs1_gbs8_mcore_te_tp4_pp1_cp1_1N8G/model_config.yaml index 7d967e68a27..23a819d08e5 100644 --- a/tests/functional_tests/test_cases/gpt-nemo/gemma2-nemo_2b_mr_mbs1_gbs8_mcore_te_tp4_pp1_cp1_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt-nemo/gemma2-nemo_2b_mr_mbs1_gbs8_mcore_te_tp4_pp1_cp1_1N8G/model_config.yaml @@ -15,3 +15,4 @@ MODEL_ARGS: data.global_batch_size: 8 data.seq_length: 2048 TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt-nemo/llama3-nemo_8b_mr_mbs1_gbs8_mcore_te_8experts_tp2_ep2_pp2_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/gpt-nemo/llama3-nemo_8b_mr_mbs1_gbs8_mcore_te_8experts_tp2_ep2_pp2_dgx_a100_1N8G/model_config.yaml index e1fb8875b56..b6d8b8647a7 100644 --- a/tests/functional_tests/test_cases/gpt-nemo/llama3-nemo_8b_mr_mbs1_gbs8_mcore_te_8experts_tp2_ep2_pp2_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt-nemo/llama3-nemo_8b_mr_mbs1_gbs8_mcore_te_8experts_tp2_ep2_pp2_dgx_a100_1N8G/model_config.yaml @@ -30,3 +30,4 @@ MODEL_ARGS: data.seq_length: 2048 log.log_dir: ${CHECKPOINT_SAVE_PATH} TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt-nemo/llama3-nemo_8b_mr_mbs4_gbs64_mcore_te_tp1_pp1_cp2_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/gpt-nemo/llama3-nemo_8b_mr_mbs4_gbs64_mcore_te_tp1_pp1_cp2_dgx_a100_1N8G/model_config.yaml index 0a7d5d8079d..fd992ef066c 100644 --- a/tests/functional_tests/test_cases/gpt-nemo/llama3-nemo_8b_mr_mbs4_gbs64_mcore_te_tp1_pp1_cp2_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt-nemo/llama3-nemo_8b_mr_mbs4_gbs64_mcore_te_tp1_pp1_cp2_dgx_a100_1N8G/model_config.yaml @@ -20,3 +20,4 @@ MODEL_ARGS: data.seq_length: 2048 log.log_dir: ${CHECKPOINT_SAVE_PATH} TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt-nemo/mixtral-nemo_8x7b_mr_mbs1_gbs8_mcore_te_tp2_pp1_ep2_1N8G/model_config.yaml b/tests/functional_tests/test_cases/gpt-nemo/mixtral-nemo_8x7b_mr_mbs1_gbs8_mcore_te_tp2_pp1_ep2_1N8G/model_config.yaml index c3dfa7845ed..5cdf82fbcd7 100644 --- a/tests/functional_tests/test_cases/gpt-nemo/mixtral-nemo_8x7b_mr_mbs1_gbs8_mcore_te_tp2_pp1_ep2_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt-nemo/mixtral-nemo_8x7b_mr_mbs1_gbs8_mcore_te_tp2_pp1_ep2_1N8G/model_config.yaml @@ -19,3 +19,4 @@ MODEL_ARGS: data.global_batch_size: 8 data.seq_length: 2048 TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt-nemo/t5-nemo_220m_mr_mbs4_gbs64_te_tp1_pp1_1N8G/model_config.yaml b/tests/functional_tests/test_cases/gpt-nemo/t5-nemo_220m_mr_mbs4_gbs64_te_tp1_pp1_1N8G/model_config.yaml index fabc337d832..e02604fe06d 100644 --- a/tests/functional_tests/test_cases/gpt-nemo/t5-nemo_220m_mr_mbs4_gbs64_te_tp1_pp1_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt-nemo/t5-nemo_220m_mr_mbs4_gbs64_te_tp1_pp1_1N8G/model_config.yaml @@ -13,3 +13,4 @@ MODEL_ARGS: data.global_batch_size: 64 data.seq_length: 512 TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_7b_tp1_pp4_memory_speed/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_7b_tp1_pp4_memory_speed/model_config.yaml index eb253b243f1..d902204b201 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_7b_tp1_pp4_memory_speed/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_7b_tp1_pp4_memory_speed/model_config.yaml @@ -67,3 +67,4 @@ METRICS: - "num-zeros" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_7b_tp4_pp1_memory_speed/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_7b_tp4_pp1_memory_speed/model_config.yaml index ae067d246c9..0dfe58f1070 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_7b_tp4_pp1_memory_speed/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_7b_tp4_pp1_memory_speed/model_config.yaml @@ -65,3 +65,4 @@ METRICS: - "num-zeros" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_disable/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_disable/model_config.yaml index 7d946d05b0c..ba949522872 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_disable/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_disable/model_config.yaml @@ -85,4 +85,5 @@ AFTER_SCRIPT: | check_log_not -F "WARNING:megatron.core.rerun_state_machine:Result validation enabled" check_log -F "Setting rerun_state_machine.current_iteration to 0..." EXIT_CODE=0 -TEST_TYPE: regular \ No newline at end of file +TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_enable/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_enable/model_config.yaml index ee557999f8e..c21ff97bbef 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_enable/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_enable/model_config.yaml @@ -82,4 +82,5 @@ AFTER_SCRIPT: | check_log() { if [[ -z $(grep -r $1 "$2" $LOG_DIR) ]]; then exit 1; else echo OK; fi } check_log -F "WARNING:megatron.core.rerun_state_machine:Result validation enabled" EXIT_CODE=0 -TEST_TYPE: regular \ No newline at end of file +TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_1/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_1/model_config.yaml index daa5093f0a4..66cc7d815b0 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_1/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_1/model_config.yaml @@ -88,3 +88,4 @@ AFTER_SCRIPT: | check_log -F "Saving a checkpoint and exiting now. Please resume the job from the checkpoint to rerun the last iteration and establish a diagnostic" EXIT_CODE=0 TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_1_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_1_1node/model_config.yaml index daa5093f0a4..66cc7d815b0 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_1_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_1_1node/model_config.yaml @@ -88,3 +88,4 @@ AFTER_SCRIPT: | check_log -F "Saving a checkpoint and exiting now. Please resume the job from the checkpoint to rerun the last iteration and establish a diagnostic" EXIT_CODE=0 TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_2/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_2/model_config.yaml index d9e8a56f0de..35a48e59081 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_2/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_persistent_2/model_config.yaml @@ -83,4 +83,5 @@ AFTER_SCRIPT: | check_log -F "WARNING:megatron.core.rerun_state_machine:Result validation enabled" check_log -E "ERROR:megatron\.core\.rerun_state_machine:Rank [0-9]+, node ([0-9a-z]|\-)+, device [0-9]+: Possible persistent error!!" EXIT_CODE=0 -TEST_TYPE: frozen-start \ No newline at end of file +TEST_TYPE: frozen-start +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_reshard/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_reshard/model_config.yaml index 7fbd18cfbe5..e229f5039bc 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_reshard/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_reshard/model_config.yaml @@ -84,4 +84,5 @@ AFTER_SCRIPT: | check_log -F "WARNING:megatron.core.rerun_state_machine:Result validation enabled" check_log -F "Job sharding has changed: Rerun state will be ignored" EXIT_CODE=0 -TEST_TYPE: frozen-start \ No newline at end of file +TEST_TYPE: frozen-start +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_resume/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_resume/model_config.yaml index f2516fd2d4b..85c2f580f27 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_resume/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_resume/model_config.yaml @@ -83,4 +83,5 @@ AFTER_SCRIPT: | check_log -F "successfully loaded checkpoint" check_log -F "WARNING:megatron.core.rerun_state_machine:Result validation enabled" EXIT_CODE=0 -TEST_TYPE: frozen-start \ No newline at end of file +TEST_TYPE: frozen-start +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_resume_check_grads/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_resume_check_grads/model_config.yaml index 1a260774210..b671f4925c1 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_resume_check_grads/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_resume_check_grads/model_config.yaml @@ -151,3 +151,4 @@ MODEL_ARGS_5: --tensor-model-parallel-size: 1 --context-parallel-size: 1 --pipeline-model-parallel-size: 2 +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_transient/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_transient/model_config.yaml index 47fd85b9d5c..097044ac1ec 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_transient/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_reruns_transient/model_config.yaml @@ -88,3 +88,4 @@ AFTER_SCRIPT: | check_log -E "ERROR:megatron\.core\.rerun_state_machine:Rank [0-9]+, node ([0-9a-z]|\-)+, device [0-9]+: Possible transient error!!" EXIT_CODE=0 TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_fim_dataset/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_fim_dataset/model_config.yaml index 86039c4d7f2..f93ab50e9c7 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_fim_dataset/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_fim_dataset/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_fim_dataset_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_fim_dataset_1node/model_config.yaml index 86039c4d7f2..f93ab50e9c7 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_fim_dataset_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_fim_dataset_1node/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_no_mmap_bin_files/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_no_mmap_bin_files/model_config.yaml index 6a90ba0f943..bacaf68d3f7 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_no_mmap_bin_files/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_no_mmap_bin_files/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_no_mmap_bin_files_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_no_mmap_bin_files_1node/model_config.yaml index 6a90ba0f943..bacaf68d3f7 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_no_mmap_bin_files_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_dist_optimizer_no_mmap_bin_files_1node/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_frozen_resume_torch_dist_dist_optimizer/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_frozen_resume_torch_dist_dist_optimizer/model_config.yaml index 8dae22b8852..cbb3a61cf83 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_frozen_resume_torch_dist_dist_optimizer/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_frozen_resume_torch_dist_dist_optimizer/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: frozen-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_mup/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_mup/model_config.yaml index ff2da3180fc..5e85ed6473c 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_mup/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_mup/model_config.yaml @@ -57,3 +57,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer/model_config.yaml index 5859b7461f9..7f3ee2a26cb 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer_1node/model_config.yaml index 5859b7461f9..7f3ee2a26cb 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer_1node/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer_no_mmap_bin_files/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer_no_mmap_bin_files/model_config.yaml index 685ec4b3db7..9a22333bba6 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer_no_mmap_bin_files/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_dist_optimizer_no_mmap_bin_files/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_uniform_full_recompute/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_uniform_full_recompute/model_config.yaml index c3ca9477dd7..2a427358771 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_uniform_full_recompute/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_uniform_full_recompute/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_uniform_full_recompute_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_uniform_full_recompute_1node/model_config.yaml index c3ca9477dd7..2a427358771 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_uniform_full_recompute_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_resume_torch_dist_uniform_full_recompute_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_uniform_full_recompute/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_uniform_full_recompute/model_config.yaml index 8d8d30fc39c..c93f4e2964d 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_uniform_full_recompute/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp1_uniform_full_recompute/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_cp4_a2a_p2p_nondeterministic/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_cp4_a2a_p2p_nondeterministic/model_config.yaml index 99de5a5d98c..3da6173b0d3 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_cp4_a2a_p2p_nondeterministic/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_cp4_a2a_p2p_nondeterministic/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_cp4_a2a_p2p_nondeterministic/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_cp4_a2a_p2p_nondeterministic/model_config.yaml index a173a0a5845..adfb883898e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_cp4_a2a_p2p_nondeterministic/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_cp4_a2a_p2p_nondeterministic/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings/model_config.yaml index 9cde3247944..ff6ca307333 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_1node/model_config.yaml index 9cde3247944..ff6ca307333 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_1node/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_interleaved_no_fusion/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_interleaved_no_fusion/model_config.yaml index b4e07b5b5e1..5c655dcfc96 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_interleaved_no_fusion/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_interleaved_no_fusion/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_interleaved_no_fusion_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_interleaved_no_fusion_1node/model_config.yaml index b4e07b5b5e1..5c655dcfc96 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_interleaved_no_fusion_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_resume_torch_dist_rope_embeddings_interleaved_no_fusion_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_rope_embeddings/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_rope_embeddings/model_config.yaml index e80995f90a3..77cd8a9758e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_rope_embeddings/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_rope_embeddings/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_rope_embeddings_interleaved_no_fusion/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_rope_embeddings_interleaved_no_fusion/model_config.yaml index 98c7db9b9bd..7499d1b6d3e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_rope_embeddings_interleaved_no_fusion/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp2_rope_embeddings_interleaved_no_fusion/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_disable_bias_linear/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_disable_bias_linear/model_config.yaml index 8d01a9132eb..dea6494695b 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_disable_bias_linear/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_disable_bias_linear/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_frozen_resume_torch_dist_swiglu/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_frozen_resume_torch_dist_swiglu/model_config.yaml index 16f9ba79fe4..9de582252a4 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_frozen_resume_torch_dist_swiglu/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_frozen_resume_torch_dist_swiglu/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: frozen-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_persistent_ckpt_disable_bias_linear/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_persistent_ckpt_disable_bias_linear/model_config.yaml index 89b883024df..d716f6c7cc0 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_persistent_ckpt_disable_bias_linear/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_persistent_ckpt_disable_bias_linear/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_disable_bias_linear/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_disable_bias_linear/model_config.yaml index 305e7eabf98..7df314462eb 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_disable_bias_linear/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_disable_bias_linear/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_disable_bias_linear_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_disable_bias_linear_1node/model_config.yaml index 305e7eabf98..7df314462eb 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_disable_bias_linear_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_disable_bias_linear_1node/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_persistent_disable_bias_linear/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_persistent_disable_bias_linear/model_config.yaml index 3cf79eaf7d2..fe0f6e9ddcb 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_persistent_disable_bias_linear/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_persistent_disable_bias_linear/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --bf16: true --log-memory-to-tensorboard: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_persistent_disable_bias_linear_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_persistent_disable_bias_linear_1node/model_config.yaml index 3cf79eaf7d2..fe0f6e9ddcb 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_persistent_disable_bias_linear_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_persistent_disable_bias_linear_1node/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --bf16: true --log-memory-to-tensorboard: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_sequence_parallel/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_sequence_parallel/model_config.yaml index 15772459af3..292bf75f976 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_sequence_parallel/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_sequence_parallel/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_swiglu/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_swiglu/model_config.yaml index a39ef7f4f78..5f8177e0eca 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_swiglu/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_swiglu/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_swiglu_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_swiglu_1node/model_config.yaml index a39ef7f4f78..5f8177e0eca 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_swiglu_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_swiglu_1node/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_untie_embeddings_and_outputs/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_untie_embeddings_and_outputs/model_config.yaml index 5d1a4257402..6d0c41792b5 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_untie_embeddings_and_outputs/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_untie_embeddings_and_outputs/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_untie_embeddings_and_outputs_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_untie_embeddings_and_outputs_1node/model_config.yaml index 5d1a4257402..6d0c41792b5 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_untie_embeddings_and_outputs_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_resume_torch_dist_untie_embeddings_and_outputs_1node/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_sequence_parallel/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_sequence_parallel/model_config.yaml index c201855b87a..d995f169bc9 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_sequence_parallel/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_sequence_parallel/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_swiglu/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_swiglu/model_config.yaml index fcadf67b1c0..8b1d7be149e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_swiglu/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_swiglu/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_untie_embeddings_and_outputs/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_untie_embeddings_and_outputs/model_config.yaml index 5b2bc318437..df7a618b6c9 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_untie_embeddings_and_outputs/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_untie_embeddings_and_outputs/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1/model_config.yaml index 8d764f5a87a..5d6feb5dde4 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_1node/model_config.yaml index 8d764f5a87a..5d6feb5dde4 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_calculate_per_token_loss/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_calculate_per_token_loss/model_config.yaml index 034339eef65..f96b0ea4686 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_calculate_per_token_loss/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_calculate_per_token_loss/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_decoupled_lr/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_decoupled_lr/model_config.yaml index bf6775690cb..f1c6025107b 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_decoupled_lr/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_decoupled_lr/model_config.yaml @@ -50,3 +50,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce/model_config.yaml index 43d9d059569..1147d997f23 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml index 153db6838cc..39d06eff08a 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml @@ -60,3 +60,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather_overlap_optimizer/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather_overlap_optimizer/model_config.yaml index 28248046b44..30d8699dc52 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather_overlap_optimizer/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather_overlap_optimizer/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather_overlap_optimizer_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather_overlap_optimizer_1node/model_config.yaml index 28248046b44..30d8699dc52 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather_overlap_optimizer_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_param_gather_overlap_optimizer_1node/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_untied/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_untied/model_config.yaml index 7c6c7ff6f9f..b6d50d2b81d 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_untied/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_dist_optimizer_overlap_grad_reduce_untied/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_decoupled_lr/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_decoupled_lr/model_config.yaml index 6bcce43a2db..d136e7afeb9 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_decoupled_lr/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_decoupled_lr/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --bf16: true --log-memory-to-tensorboard: true TEST_TYPE: regular # Usually ckpt-resume, but as a WAR to #513 set to regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_decoupled_lr_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_decoupled_lr_1node/model_config.yaml index 6bcce43a2db..d136e7afeb9 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_decoupled_lr_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_decoupled_lr_1node/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --bf16: true --log-memory-to-tensorboard: true TEST_TYPE: regular # Usually ckpt-resume, but as a WAR to #513 set to regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist/model_config.yaml index 51cbf48c21b..a5347336e67 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular # Usually ckpt-resume, but as a WAR to #513 set to regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_calculate_per_token_loss/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_calculate_per_token_loss/model_config.yaml index 6d77c3df361..d1e0bb6b16e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_calculate_per_token_loss/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_calculate_per_token_loss/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_calculate_per_token_loss_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_calculate_per_token_loss_1node/model_config.yaml index 6d77c3df361..d1e0bb6b16e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_calculate_per_token_loss_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_calculate_per_token_loss_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce/model_config.yaml index 0a1f510fc8f..415acdf4717 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_1node/model_config.yaml index 0a1f510fc8f..415acdf4717 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_untied/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_untied/model_config.yaml index a8194368c4e..e929bba269f 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_untied/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_untied/model_config.yaml @@ -58,3 +58,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular # Usually ckpt-resume, but as a WAR to #513 set to regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_tunable_overlap/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_tunable_overlap/model_config.yaml index 4dfc41a948f..f78a4fd10bc 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_tunable_overlap/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_resume_torch_dist_tunable_overlap/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular # Usually ckpt-resume, but as a WAR to #513 set to regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_tunable_overlap/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_tunable_overlap/model_config.yaml index 7dcdd335e8c..a87c869dcd1 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_tunable_overlap/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_tunable_overlap/model_config.yaml @@ -57,3 +57,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_tunable_overlap_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_tunable_overlap_1node/model_config.yaml index 7dcdd335e8c..a87c869dcd1 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_tunable_overlap_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_tunable_overlap_1node/model_config.yaml @@ -57,3 +57,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_uneven_pipeline/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_uneven_pipeline/model_config.yaml index e068c864d10..7ae45f52dbf 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_uneven_pipeline/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_uneven_pipeline/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_uneven_pipeline_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_uneven_pipeline_1node/model_config.yaml index e068c864d10..7ae45f52dbf 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_uneven_pipeline_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp1_uneven_pipeline_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp2_account_for_embedding_loss_in_pipeline_split/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp2_account_for_embedding_loss_in_pipeline_split/model_config.yaml index b9e66c55bff..798224308d4 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp2_account_for_embedding_loss_in_pipeline_split/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp2_account_for_embedding_loss_in_pipeline_split/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp2_account_for_embedding_loss_in_pipeline_split_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp2_account_for_embedding_loss_in_pipeline_split_1node/model_config.yaml index b9e66c55bff..798224308d4 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp2_account_for_embedding_loss_in_pipeline_split_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp1_pp4_vp2_account_for_embedding_loss_in_pipeline_split_1node/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_cp2_nondeterministic/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_cp2_nondeterministic/model_config.yaml index 2f33dea359d..b798aa016e1 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_cp2_nondeterministic/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_cp2_nondeterministic/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_cp2_nondeterministic/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_cp2_nondeterministic/model_config.yaml index 0d803ff1aaf..f2affd7e792 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_cp2_nondeterministic/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_cp2_nondeterministic/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: frozen-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_fsdp2_resume_torch_dist/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_fsdp2_resume_torch_dist/model_config.yaml index 9876606f2f7..2a52d3c1250 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_fsdp2_resume_torch_dist/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_fsdp2_resume_torch_dist/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn/model_config.yaml index cc7211c9967..4127f72642b 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn/model_config.yaml @@ -80,3 +80,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_1node/model_config.yaml index cc7211c9967..4127f72642b 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_1node/model_config.yaml @@ -80,3 +80,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_async/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_async/model_config.yaml index 620b4c2b734..48ba564b1f8 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_async/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_async/model_config.yaml @@ -85,4 +85,5 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: regular -TEST_EVALUATION: xpass \ No newline at end of file +TEST_EVALUATION: xpass +LAUNCHER: torchrun diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_async_mcore/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_async_mcore/model_config.yaml index 731a2927b22..4d361a20ce6 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_async_mcore/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_async_mcore/model_config.yaml @@ -85,4 +85,5 @@ MODEL_ARGS: --use-persistent-ckpt-worker: true --attention-backend: unfused --log-memory-to-tensorboard: true -TEST_TYPE: regular \ No newline at end of file +TEST_TYPE: regular +LAUNCHER: torchrun diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_sync/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_sync/model_config.yaml index 6dfd675c343..b354332b027 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_sync/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_gdn_no_nvrx_sync/model_config.yaml @@ -82,4 +82,5 @@ MODEL_ARGS: --bf16: true --attention-backend: unfused --log-memory-to-tensorboard: true -TEST_TYPE: regular \ No newline at end of file +TEST_TYPE: regular +LAUNCHER: torchrun diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_dev_dgx_a100.json b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_dev_dgx_a100.json deleted file mode 100644 index ae531bf007e..00000000000 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_dev_dgx_a100.json +++ /dev/null @@ -1,109 +0,0 @@ -{ - "kd loss": { - "start_step": 1, - "end_step": 100, - "step_interval": 1, - "values": { - "1": 0.4930937, - "2": 0.4935429, - "3": 0.4938237, - "4": 0.4938229, - "5": 0.4924306, - "6": 0.4935288, - "7": 0.4928354, - "8": 0.4925097, - "9": 0.4936136, - "10": 0.4922911, - "11": 0.4934031, - "12": 0.4951033, - "13": 0.4918853, - "14": 0.4936183, - "15": 0.4926639, - "16": 0.4927304, - "17": 0.4925308, - "18": 0.4927951, - "19": 0.4938825, - "20": 0.4939776, - "21": 0.4933512, - "22": 0.4935322, - "23": 0.4937269, - "24": 0.4927326, - "25": 0.4927868, - "26": 0.4927689, - "27": 0.4924214, - "28": 0.4925573, - "29": 0.4917694, - "30": 0.4919884, - "31": 0.4929765, - "32": 0.4930308, - "33": 0.4928029, - "34": 0.4923102, - "35": 0.4918847, - "36": 0.4914086, - "37": 0.4929215, - "38": 0.4923307, - "39": 0.4910690, - "40": 0.4919418, - "41": 0.4913271, - "42": 0.4919568, - "43": 0.4903573, - "44": 0.4916522, - "45": 0.4915655, - "46": 0.4898856, - "47": 0.4899229, - "48": 0.4892673, - "49": 0.4894423, - "50": 0.4903796, - "51": 0.4907262, - "52": 0.4882944, - "53": 0.4877340, - "54": 0.4902404, - "55": 0.4881638, - "56": 0.4888564, - "57": 0.4882180, - "58": 0.4887677, - "59": 0.4883497, - "60": 0.4863744, - "61": 0.4875762, - "62": 0.4837778, - "63": 0.4867221, - "64": 0.4840697, - "65": 0.4840384, - "66": 0.4857976, - "67": 0.4837634, - "68": 0.4800620, - "69": 0.4781690, - "70": 0.4818793, - "71": 0.4796092, - "72": 0.4783594, - "73": 0.4789546, - "74": 0.4767389, - "75": 0.4774750, - "76": 0.4746155, - "77": 0.4745574, - "78": 0.4737080, - "79": 0.4718909, - "80": 0.4693059, - "81": 0.4696763, - "82": 0.4705839, - "83": 0.4661151, - "84": 0.4634311, - "85": 0.4654806, - "86": 0.4634389, - "87": 0.4595202, - "88": 0.4586285, - "89": 0.4595589, - "90": 0.4561397, - "91": 0.4547178, - "92": 0.4553441, - "93": 0.4545716, - "94": 0.4518545, - "95": 0.4504384, - "96": 0.4518549, - "97": 0.4450935, - "98": 0.4441919, - "99": 0.4442465, - "100": 0.4427010 - } - } -} diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_dev_dgx_gb200.json b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_dev_dgx_gb200.json new file mode 100644 index 00000000000..68a37fce071 --- /dev/null +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_dev_dgx_gb200.json @@ -0,0 +1,643 @@ +{ + "lm loss": { + "start_step": 1, + "end_step": 100, + "step_interval": 1, + "values": { + "1": 0.0, + "2": 0.0, + "3": 0.0, + "4": 0.0, + "5": 0.0, + "6": 0.0, + "7": 0.0, + "8": 0.0, + "9": 0.0, + "10": 0.0, + "11": 0.0, + "12": 0.0, + "13": 0.0, + "14": 0.0, + "15": 0.0, + "16": 0.0, + "17": 0.0, + "18": 0.0, + "19": 0.0, + "20": 0.0, + "21": 0.0, + "22": 0.0, + "23": 0.0, + "24": 0.0, + "25": 0.0, + "26": 0.0, + "27": 0.0, + "28": 0.0, + "29": 0.0, + "30": 0.0, + "31": 0.0, + "32": 0.0, + "33": 0.0, + "34": 0.0, + "35": 0.0, + "36": 0.0, + "37": 0.0, + "38": 0.0, + "39": 0.0, + "40": 0.0, + "41": 0.0, + "42": 0.0, + "43": 0.0, + "44": 0.0, + "45": 0.0, + "46": 0.0, + "47": 0.0, + "48": 0.0, + "49": 0.0, + "50": 0.0, + "51": 0.0, + "52": 0.0, + "53": 0.0, + "54": 0.0, + "55": 0.0, + "56": 0.0, + "57": 0.0, + "58": 0.0, + "59": 0.0, + "60": 0.0, + "61": 0.0, + "62": 0.0, + "63": 0.0, + "64": 0.0, + "65": 0.0, + "66": 0.0, + "67": 0.0, + "68": 0.0, + "69": 0.0, + "70": 0.0, + "71": 0.0, + "72": 0.0, + "73": 0.0, + "74": 0.0, + "75": 0.0, + "76": 0.0, + "77": 0.0, + "78": 0.0, + "79": 0.0, + "80": 0.0, + "81": 0.0, + "82": 0.0, + "83": 0.0, + "84": 0.0, + "85": 0.0, + "86": 0.0, + "87": 0.0, + "88": 0.0, + "89": 0.0, + "90": 0.0, + "91": 0.0, + "92": 0.0, + "93": 0.0, + "94": 0.0, + "95": 0.0, + "96": 0.0, + "97": 0.0, + "98": 0.0, + "99": 0.0, + "100": 0.0 + } + }, + "total loss": { + "start_step": 1, + "end_step": 100, + "step_interval": 1, + "values": { + "1": 0.49229, + "2": 0.49235, + "3": 0.49306, + "4": 0.49235, + "5": 0.49191, + "6": 0.49253, + "7": 0.49213, + "8": 0.49258, + "9": 0.4937, + "10": 0.4928, + "11": 0.49252, + "12": 0.49348, + "13": 0.49141, + "14": 0.49211, + "15": 0.49198, + "16": 0.49259, + "17": 0.49109, + "18": 0.49167, + "19": 0.49252, + "20": 0.49093, + "21": 0.4912, + "22": 0.49168, + "23": 0.49231, + "24": 0.49276, + "25": 0.49294, + "26": 0.49277, + "27": 0.49224, + "28": 0.49234, + "29": 0.49234, + "30": 0.49298, + "31": 0.49185, + "32": 0.49196, + "33": 0.49228, + "34": 0.49226, + "35": 0.49223, + "36": 0.49134, + "37": 0.49055, + "38": 0.4907, + "39": 0.49142, + "40": 0.49151, + "41": 0.49131, + "42": 0.49104, + "43": 0.49075, + "44": 0.49093, + "45": 0.4897, + "46": 0.49004, + "47": 0.49013, + "48": 0.48992, + "49": 0.48896, + "50": 0.48992, + "51": 0.48905, + "52": 0.48885, + "53": 0.49017, + "54": 0.48791, + "55": 0.48706, + "56": 0.48701, + "57": 0.48685, + "58": 0.48532, + "59": 0.48558, + "60": 0.48393, + "61": 0.48541, + "62": 0.48361, + "63": 0.48138, + "64": 0.48017, + "65": 0.47951, + "66": 0.47917, + "67": 0.47919, + "68": 0.47884, + "69": 0.47925, + "70": 0.47516, + "71": 0.47456, + "72": 0.47368, + "73": 0.473, + "74": 0.47029, + "75": 0.47117, + "76": 0.46983, + "77": 0.46684, + "78": 0.46518, + "79": 0.46566, + "80": 0.46487, + "81": 0.46125, + "82": 0.46293, + "83": 0.45894, + "84": 0.45692, + "85": 0.45714, + "86": 0.45544, + "87": 0.45196, + "88": 0.45104, + "89": 0.45379, + "90": 0.44862, + "91": 0.4482, + "92": 0.4449, + "93": 0.44739, + "94": 0.4421, + "95": 0.44338, + "96": 0.44316, + "97": 0.43975, + "98": 0.44261, + "99": 0.43686, + "100": 0.43249 + } + }, + "num-zeros": { + "start_step": 1, + "end_step": 100, + "step_interval": 1, + "values": { + "1": 24026928.0, + "2": 24143812.0, + "3": 24187244.0, + "4": 23877068.0, + "5": 23864828.0, + "6": 24045444.0, + "7": 23949008.0, + "8": 24026020.0, + "9": 23886348.0, + "10": 24120268.0, + "11": 24066420.0, + "12": 23899156.0, + "13": 23975336.0, + "14": 24178010.0, + "15": 23831062.0, + "16": 23945590.0, + "17": 23816820.0, + "18": 24135244.0, + "19": 23732252.0, + "20": 23890840.0, + "21": 23832704.0, + "22": 24119708.0, + "23": 23826508.0, + "24": 23976604.0, + "25": 24147570.0, + "26": 24493544.0, + "27": 24200086.0, + "28": 24034610.0, + "29": 23844336.0, + "30": 24090330.0, + "31": 23953480.0, + "32": 24038172.0, + "33": 24007426.0, + "34": 24113576.0, + "35": 24080250.0, + "36": 24163534.0, + "37": 24106944.0, + "38": 24232868.0, + "39": 23929552.0, + "40": 24008410.0, + "41": 23998468.0, + "42": 24167966.0, + "43": 24092170.0, + "44": 24091696.0, + "45": 24073596.0, + "46": 23743600.0, + "47": 24161662.0, + "48": 23921924.0, + "49": 24004830.0, + "50": 23789460.0, + "51": 24147816.0, + "52": 24110948.0, + "53": 23841400.0, + "54": 24148992.0, + "55": 24145062.0, + "56": 24073800.0, + "57": 24208330.0, + "58": 24028408.0, + "59": 24173804.0, + "60": 23985040.0, + "61": 24130440.0, + "62": 24025044.0, + "63": 24060980.0, + "64": 23777584.0, + "65": 23728788.0, + "66": 24203238.0, + "67": 24283560.0, + "68": 24057064.0, + "69": 24046326.0, + "70": 23915612.0, + "71": 23818212.0, + "72": 24018984.0, + "73": 24159400.0, + "74": 24251924.0, + "75": 24123716.0, + "76": 24075268.0, + "77": 24195152.0, + "78": 23979564.0, + "79": 23984260.0, + "80": 23956760.0, + "81": 24113864.0, + "82": 24171568.0, + "83": 24025912.0, + "84": 23832644.0, + "85": 24063532.0, + "86": 24206276.0, + "87": 23923704.0, + "88": 23940504.0, + "89": 24254836.0, + "90": 23856988.0, + "91": 23983564.0, + "92": 23931634.0, + "93": 24104520.0, + "94": 23955112.0, + "95": 24049972.0, + "96": 24021030.0, + "97": 23738292.0, + "98": 24436300.0, + "99": 23954008.0, + "100": 23844172.0 + } + }, + "mem-allocated-bytes": { + "start_step": 1, + "end_step": 100, + "step_interval": 1, + "values": { + "1": 733251072.0, + "2": 733251072.0, + "3": 733251072.0, + "4": 733251072.0, + "5": 733251072.0, + "6": 733251072.0, + "7": 733251072.0, + "8": 733251072.0, + "9": 733251072.0, + "10": 733251072.0, + "11": 733251072.0, + "12": 733251072.0, + "13": 733251072.0, + "14": 733251072.0, + "15": 733251072.0, + "16": 733251072.0, + "17": 733251072.0, + "18": 733251072.0, + "19": 733251072.0, + "20": 733251072.0, + "21": 733251072.0, + "22": 733251072.0, + "23": 733251072.0, + "24": 733251072.0, + "25": 733251072.0, + "26": 733251072.0, + "27": 733251072.0, + "28": 733251072.0, + "29": 733251072.0, + "30": 733251072.0, + "31": 733251072.0, + "32": 733251072.0, + "33": 733251072.0, + "34": 733251072.0, + "35": 733251072.0, + "36": 733251072.0, + "37": 733251072.0, + "38": 733251072.0, + "39": 733251072.0, + "40": 733251072.0, + "41": 733251072.0, + "42": 733251072.0, + "43": 733251072.0, + "44": 733251072.0, + "45": 733251072.0, + "46": 733251072.0, + "47": 733251072.0, + "48": 733251072.0, + "49": 733251072.0, + "50": 733251072.0, + "51": 733251072.0, + "52": 733251072.0, + "53": 733251072.0, + "54": 733251072.0, + "55": 733251072.0, + "56": 733251072.0, + "57": 733251072.0, + "58": 733251072.0, + "59": 733251072.0, + "60": 733251072.0, + "61": 733251072.0, + "62": 733251072.0, + "63": 733251072.0, + "64": 733251072.0, + "65": 733251072.0, + "66": 733251072.0, + "67": 733251072.0, + "68": 733251072.0, + "69": 733251072.0, + "70": 733251072.0, + "71": 733251072.0, + "72": 733251072.0, + "73": 733251072.0, + "74": 733251072.0, + "75": 733251072.0, + "76": 733251072.0, + "77": 733251072.0, + "78": 733251072.0, + "79": 733251072.0, + "80": 733251072.0, + "81": 733251072.0, + "82": 733251072.0, + "83": 733251072.0, + "84": 733251072.0, + "85": 733251072.0, + "86": 733251072.0, + "87": 733251072.0, + "88": 733251072.0, + "89": 733251072.0, + "90": 733251072.0, + "91": 733251072.0, + "92": 733251072.0, + "93": 733251072.0, + "94": 733251072.0, + "95": 733251072.0, + "96": 733251072.0, + "97": 733251072.0, + "98": 733251072.0, + "99": 733251072.0, + "100": 733251072.0 + } + }, + "mem-max-allocated-bytes": { + "start_step": 1, + "end_step": 100, + "step_interval": 1, + "values": { + "1": 3237344256.0, + "2": 3326439424.0, + "3": 3326439424.0, + "4": 3326439424.0, + "5": 3326439424.0, + "6": 3326439424.0, + "7": 3326439424.0, + "8": 3326439424.0, + "9": 3326439424.0, + "10": 3326439424.0, + "11": 3326439424.0, + "12": 3326439424.0, + "13": 3326439424.0, + "14": 3326439424.0, + "15": 3326504960.0, + "16": 3326504960.0, + "17": 3326504960.0, + "18": 3326504960.0, + "19": 3326504960.0, + "20": 3326504960.0, + "21": 3326504960.0, + "22": 3326504960.0, + "23": 3326504960.0, + "24": 3326504960.0, + "25": 3326504960.0, + "26": 3326504960.0, + "27": 3326504960.0, + "28": 3326504960.0, + "29": 3326504960.0, + "30": 3326504960.0, + "31": 3326504960.0, + "32": 3326504960.0, + "33": 3326504960.0, + "34": 3326504960.0, + "35": 3326504960.0, + "36": 3326963712.0, + "37": 3326963712.0, + "38": 3326963712.0, + "39": 3326963712.0, + "40": 3326963712.0, + "41": 3326963712.0, + "42": 3326963712.0, + "43": 3326963712.0, + "44": 3326963712.0, + "45": 3326963712.0, + "46": 3326963712.0, + "47": 3326963712.0, + "48": 3326963712.0, + "49": 3326963712.0, + "50": 3326963712.0, + "51": 3328271872.0, + "52": 3328271872.0, + "53": 3328271872.0, + "54": 3328271872.0, + "55": 3328271872.0, + "56": 3328271872.0, + "57": 3328271872.0, + "58": 3328271872.0, + "59": 3328271872.0, + "60": 3328271872.0, + "61": 3328271872.0, + "62": 3328271872.0, + "63": 3328271872.0, + "64": 3328271872.0, + "65": 3328271872.0, + "66": 3328271872.0, + "67": 3328271872.0, + "68": 3328271872.0, + "69": 3328271872.0, + "70": 3328271872.0, + "71": 3328271872.0, + "72": 3328271872.0, + "73": 3328271872.0, + "74": 3328271872.0, + "75": 3328271872.0, + "76": 3328271872.0, + "77": 3328271872.0, + "78": 3328271872.0, + "79": 3328271872.0, + "80": 3328271872.0, + "81": 3328271872.0, + "82": 3328271872.0, + "83": 3328271872.0, + "84": 3328271872.0, + "85": 3328271872.0, + "86": 3328271872.0, + "87": 3328271872.0, + "88": 3328271872.0, + "89": 3328271872.0, + "90": 3328271872.0, + "91": 3328271872.0, + "92": 3328271872.0, + "93": 3328271872.0, + "94": 3328271872.0, + "95": 3328271872.0, + "96": 3328271872.0, + "97": 3328271872.0, + "98": 3328271872.0, + "99": 3328271872.0, + "100": 3328271872.0 + } + }, + "iteration-time": { + "start_step": 2, + "end_step": 100, + "step_interval": 1, + "values": { + "2": 5.4207, + "3": 0.32071, + "4": 0.32419, + "5": 0.27463, + "6": 0.27557, + "7": 0.27601, + "8": 0.27318, + "9": 0.27411, + "10": 0.27381, + "11": 0.27545, + "12": 0.30916, + "13": 0.31371, + "14": 0.3026, + "15": 0.29432, + "16": 0.31261, + "17": 0.31431, + "18": 0.30656, + "19": 0.3174, + "20": 0.32061, + "21": 0.31778, + "22": 0.35357, + "23": 0.3637, + "24": 0.40747, + "25": 0.37189, + "26": 0.40857, + "27": 0.30615, + "28": 0.31235, + "29": 0.31541, + "30": 0.4272, + "31": 0.40504, + "32": 0.51821, + "33": 0.46673, + "34": 0.42968, + "35": 0.414, + "36": 0.39324, + "37": 0.31311, + "38": 0.30297, + "39": 0.3022, + "40": 0.30204, + "41": 0.30404, + "42": 0.3041, + "43": 0.31703, + "44": 0.30981, + "45": 0.30867, + "46": 0.30996, + "47": 0.31229, + "48": 0.31038, + "49": 0.30472, + "50": 0.32046, + "51": 0.57524, + "52": 0.36905, + "53": 0.3012, + "54": 0.31617, + "55": 0.31476, + "56": 0.32079, + "57": 0.31553, + "58": 0.30275, + "59": 0.33016, + "60": 0.32202, + "61": 0.31778, + "62": 0.32418, + "63": 0.31518, + "64": 0.3137, + "65": 0.31372, + "66": 0.31738, + "67": 0.31707, + "68": 0.31289, + "69": 0.31334, + "70": 0.311, + "71": 0.30529, + "72": 0.31341, + "73": 0.30623, + "74": 0.31657, + "75": 0.31992, + "76": 0.31396, + "77": 0.31781, + "78": 0.32836, + "79": 0.32318, + "80": 0.32087, + "81": 0.31407, + "82": 0.30679, + "83": 0.3123, + "84": 0.31429, + "85": 0.31233, + "86": 0.30659, + "87": 0.31428, + "88": 0.32007, + "89": 0.34521, + "90": 0.32648, + "91": 0.31607, + "92": 0.3249, + "93": 0.32018, + "94": 0.32003, + "95": 0.31498, + "96": 0.30649, + "97": 0.30738, + "98": 0.30612, + "99": 0.31249, + "100": 0.30688 + } + } +} \ No newline at end of file diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_dev_dgx_h100.json new file mode 100644 index 00000000000..8634a66847a --- /dev/null +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_dev_dgx_h100.json @@ -0,0 +1,216 @@ +{ + "total loss": { + "start_step": 1, + "end_step": 100, + "step_interval": 1, + "values": { + "1": 0.49197, + "2": 0.49244, + "3": 0.49244, + "4": 0.49309, + "5": 0.49199, + "6": 0.4925, + "7": 0.49298, + "8": 0.49175, + "9": 0.4929, + "10": 0.49224, + "11": 0.49354, + "12": 0.49217, + "13": 0.4921, + "14": 0.49196, + "15": 0.49207, + "16": 0.49239, + "17": 0.49231, + "18": 0.49241, + "19": 0.49228, + "20": 0.49096, + "21": 0.49281, + "22": 0.49273, + "23": 0.49169, + "24": 0.49298, + "25": 0.49222, + "26": 0.49219, + "27": 0.49351, + "28": 0.4928, + "29": 0.49313, + "30": 0.49276, + "31": 0.49254, + "32": 0.49177, + "33": 0.49254, + "34": 0.49255, + "35": 0.49246, + "36": 0.49135, + "37": 0.49143, + "38": 0.49129, + "39": 0.49205, + "40": 0.49215, + "41": 0.49167, + "42": 0.49202, + "43": 0.49159, + "44": 0.49097, + "45": 0.49001, + "46": 0.49096, + "47": 0.49014, + "48": 0.4894, + "49": 0.48931, + "50": 0.48995, + "51": 0.49027, + "52": 0.48921, + "53": 0.4916, + "54": 0.49015, + "55": 0.4892, + "56": 0.48764, + "57": 0.48865, + "58": 0.4877, + "59": 0.4865, + "60": 0.48435, + "61": 0.48678, + "62": 0.48624, + "63": 0.48259, + "64": 0.48274, + "65": 0.48095, + "66": 0.48127, + "67": 0.48124, + "68": 0.47975, + "69": 0.47882, + "70": 0.47826, + "71": 0.47797, + "72": 0.47728, + "73": 0.47533, + "74": 0.47329, + "75": 0.47452, + "76": 0.4729, + "77": 0.47196, + "78": 0.46773, + "79": 0.46857, + "80": 0.46752, + "81": 0.46539, + "82": 0.46683, + "83": 0.46365, + "84": 0.45854, + "85": 0.45937, + "86": 0.46228, + "87": 0.45535, + "88": 0.45485, + "89": 0.45797, + "90": 0.44956, + "91": 0.45188, + "92": 0.44878, + "93": 0.4514, + "94": 0.44644, + "95": 0.44778, + "96": 0.44672, + "97": 0.44107, + "98": 0.44625, + "99": 0.43963, + "100": 0.43588 + } + }, + "lm loss": { + "start_step": 1, + "end_step": 100, + "step_interval": 1, + "values": { + "1": 0.0, + "2": 0.0, + "3": 0.0, + "4": 0.0, + "5": 0.0, + "6": 0.0, + "7": 0.0, + "8": 0.0, + "9": 0.0, + "10": 0.0, + "11": 0.0, + "12": 0.0, + "13": 0.0, + "14": 0.0, + "15": 0.0, + "16": 0.0, + "17": 0.0, + "18": 0.0, + "19": 0.0, + "20": 0.0, + "21": 0.0, + "22": 0.0, + "23": 0.0, + "24": 0.0, + "25": 0.0, + "26": 0.0, + "27": 0.0, + "28": 0.0, + "29": 0.0, + "30": 0.0, + "31": 0.0, + "32": 0.0, + "33": 0.0, + "34": 0.0, + "35": 0.0, + "36": 0.0, + "37": 0.0, + "38": 0.0, + "39": 0.0, + "40": 0.0, + "41": 0.0, + "42": 0.0, + "43": 0.0, + "44": 0.0, + "45": 0.0, + "46": 0.0, + "47": 0.0, + "48": 0.0, + "49": 0.0, + "50": 0.0, + "51": 0.0, + "52": 0.0, + "53": 0.0, + "54": 0.0, + "55": 0.0, + "56": 0.0, + "57": 0.0, + "58": 0.0, + "59": 0.0, + "60": 0.0, + "61": 0.0, + "62": 0.0, + "63": 0.0, + "64": 0.0, + "65": 0.0, + "66": 0.0, + "67": 0.0, + "68": 0.0, + "69": 0.0, + "70": 0.0, + "71": 0.0, + "72": 0.0, + "73": 0.0, + "74": 0.0, + "75": 0.0, + "76": 0.0, + "77": 0.0, + "78": 0.0, + "79": 0.0, + "80": 0.0, + "81": 0.0, + "82": 0.0, + "83": 0.0, + "84": 0.0, + "85": 0.0, + "86": 0.0, + "87": 0.0, + "88": 0.0, + "89": 0.0, + "90": 0.0, + "91": 0.0, + "92": 0.0, + "93": 0.0, + "94": 0.0, + "95": 0.0, + "96": 0.0, + "97": 0.0, + "98": 0.0, + "99": 0.0, + "100": 0.0 + } + } +} diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_lts_dgx_a100.json b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_lts_dgx_a100.json deleted file mode 100644 index 9e26dfeeb6e..00000000000 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/golden_values_lts_dgx_a100.json +++ /dev/null @@ -1 +0,0 @@ -{} \ No newline at end of file diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/model_config.yaml index c75a5a81414..9576723f116 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume/model_config.yaml @@ -1,11 +1,10 @@ ENV_VARS: - SKIP_PYTEST: 1 CUDA_DEVICE_MAX_CONNECTIONS: 1 NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 NCCL_ALGO: Ring CUBLAS_WORKSPACE_CONFIG: :4096:8 ARTIFACTS_ROOT: /workspace/checkpoints - DISTILL_CONFIG: '{intermediate_layer_pairs: [["decoder.final_layernorm", "decoder.final_layernorm"]], logit_layers: ["output_layer", "output_layer"], skip_lm_loss: true, kd_loss_scale: 10.0}' + DISTILL_CONFIG: '{intermediate_layer_pairs: [["decoder.final_layernorm", "decoder.final_layernorm"]], logit_layers: ["output_layer", "output_layer"], skip_lm_loss: true, kd_loss_scale: 1.0}' BEFORE_SCRIPT: | mkdir -p ${DATA_CACHE_PATH}/distill && echo $DISTILL_CONFIG | yq -P > ${DATA_CACHE_PATH}/distill/distill_config.yaml MODEL_ARGS: @@ -68,4 +67,9 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --async-save: true --use-persistent-ckpt-worker: true + --exit-interval: 100 TEST_TYPE: ckpt-resume +METRICS: + - lm loss + - total loss +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume_1node/model_config.yaml index c75a5a81414..494b685d13a 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_modelopt_distill_resume_1node/model_config.yaml @@ -69,3 +69,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_multi_dist_optimizer_instances/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_multi_dist_optimizer_instances/model_config.yaml index e541610ae72..9a7cd250f08 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_multi_dist_optimizer_instances/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_multi_dist_optimizer_instances/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_resume_torch_dist_cp2_nondeterministic/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_resume_torch_dist_cp2_nondeterministic/model_config.yaml index bd303906bd6..d3a2debce80 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_resume_torch_dist_cp2_nondeterministic/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_resume_torch_dist_cp2_nondeterministic/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_resume_torch_dist_cp2_nondeterministic_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_resume_torch_dist_cp2_nondeterministic_1node/model_config.yaml index bd303906bd6..d3a2debce80 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_resume_torch_dist_cp2_nondeterministic_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp1_resume_torch_dist_cp2_nondeterministic_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2/model_config.yaml index 11b6f1c57a1..ecdeafbc27c 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2/model_config.yaml index fc6d56fab55..45ea1cd9483 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_1node/model_config.yaml index 7be86a80447..ffd6294c4bd 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_1node/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss/model_config.yaml index 4aa6deabd64..841ffb86163 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_1node/model_config.yaml index 27fd45b5701..ede6949b68d 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_1node/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_nondeterministic/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_nondeterministic/model_config.yaml index f1d41d4a22a..cf00d693fe8 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_nondeterministic/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_nondeterministic/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_nondeterministic_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_nondeterministic_1node/model_config.yaml index 27ff315894a..2044b3b9011 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_nondeterministic_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_calculate_per_token_loss_nondeterministic_1node/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_dp_last/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_dp_last/model_config.yaml index eb6ae776b50..9a0f0a12912 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_dp_last/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_dp_last/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_dp_last_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_dp_last_1node/model_config.yaml index caf364efb28..e88fecd2f87 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_dp_last_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_dp_last_1node/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_nondeterministic_dp_last/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_nondeterministic_dp_last/model_config.yaml index 98c4aefea36..a6268baefc2 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_nondeterministic_dp_last/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_nondeterministic_dp_last/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_nondeterministic_dp_last_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_nondeterministic_dp_last_1node/model_config.yaml index c1087788693..951ea39e0e3 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_nondeterministic_dp_last_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_calculate_per_token_loss_nondeterministic_dp_last_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_dp_last/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_dp_last/model_config.yaml index f798a99d703..93148b90ef5 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_dp_last/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_dp_last/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_dp_last_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_dp_last_1node/model_config.yaml index 67831941bf5..a9f579fc051 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_dp_last_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_dp_last_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_nondeterministic_dp_last/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_nondeterministic_dp_last/model_config.yaml index 4b8aad6f105..8279ca3ca04 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_nondeterministic_dp_last/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_nondeterministic_dp_last/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_nondeterministic_dp_last_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_nondeterministic_dp_last_1node/model_config.yaml index 72faa6ac2da..b671f0ef8ba 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_nondeterministic_dp_last_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_etp4_nondeterministic_dp_last_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_nondeterministic/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_nondeterministic/model_config.yaml index 0216d5283d7..777aebe3f57 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_nondeterministic/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_nondeterministic/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_nondeterministic_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_nondeterministic_1node/model_config.yaml index 2f33dea359d..b798aa016e1 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_nondeterministic_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cp2_nondeterministic_1node/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cross_entropy_loss_fusion/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cross_entropy_loss_fusion/model_config.yaml index 2f73bf24f01..c0633f3359e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cross_entropy_loss_fusion/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cross_entropy_loss_fusion/model_config.yaml @@ -49,3 +49,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cross_entropy_loss_fusion_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cross_entropy_loss_fusion_1node/model_config.yaml index 2f73bf24f01..c0633f3359e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cross_entropy_loss_fusion_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_cross_entropy_loss_fusion_1node/model_config.yaml @@ -49,3 +49,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_ddp_average_in_collective/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_ddp_average_in_collective/model_config.yaml index 4eb20fa0683..dddef126f20 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_ddp_average_in_collective/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_ddp_average_in_collective/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_defer_embedding_wgrad_compute/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_defer_embedding_wgrad_compute/model_config.yaml index bc37d8cd771..193c8e6e144 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_defer_embedding_wgrad_compute/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_defer_embedding_wgrad_compute/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml index 507e8de9df7..6dcd6b02eb3 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_dsa/model_config.yaml @@ -64,3 +64,4 @@ MODEL_ARGS: --bf16: true --log-memory-to-tensorboard: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_mla/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_mla/model_config.yaml index d8afdbf0756..110bcf8bae8 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_mla/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_mla/model_config.yaml @@ -59,3 +59,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_mla_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_mla_1node/model_config.yaml index d8afdbf0756..110bcf8bae8 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_mla_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_mla_1node/model_config.yaml @@ -59,3 +59,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_no_create_attention_mask_in_dataloader/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_no_create_attention_mask_in_dataloader/model_config.yaml index 023fc7b6e5b..95bf40ac648 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_no_create_attention_mask_in_dataloader/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_no_create_attention_mask_in_dataloader/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_no_mmap_bin_files/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_no_mmap_bin_files/model_config.yaml index 9e0518f9ffd..1368abfc7b6 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_no_mmap_bin_files/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_no_mmap_bin_files/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist/model_config.yaml index 14e40a430d5..5ef5ede58a5 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --use-persistent-ckpt-worker: true --verify-integrity: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_1node/model_config.yaml index 9436fa2a5e6..c0d19c140b2 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_1node/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cp2_nondeterministic/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cp2_nondeterministic/model_config.yaml index 29a9bbef0c1..f81a6a59243 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cp2_nondeterministic/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cp2_nondeterministic/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --use-persistent-ckpt-worker: true --verify-integrity: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cp2_nondeterministic_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cp2_nondeterministic_1node/model_config.yaml index c682b78eedc..b1a5757076a 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cp2_nondeterministic_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cp2_nondeterministic_1node/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cross_entropy_loss_fusion/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cross_entropy_loss_fusion/model_config.yaml index 77789c90192..40d83402a4c 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cross_entropy_loss_fusion/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cross_entropy_loss_fusion/model_config.yaml @@ -51,3 +51,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cross_entropy_loss_fusion_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cross_entropy_loss_fusion_1node/model_config.yaml index 77789c90192..40d83402a4c 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cross_entropy_loss_fusion_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_cross_entropy_loss_fusion_1node/model_config.yaml @@ -51,3 +51,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_ddp_average_in_collective/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_ddp_average_in_collective/model_config.yaml index aacc077e937..0f4d863c12c 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_ddp_average_in_collective/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_ddp_average_in_collective/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_ddp_average_in_collective_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_ddp_average_in_collective_1node/model_config.yaml index aacc077e937..0f4d863c12c 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_ddp_average_in_collective_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_ddp_average_in_collective_1node/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_defer_embedding_wgrad_compute/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_defer_embedding_wgrad_compute/model_config.yaml index f600f587ae9..340f6125883 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_defer_embedding_wgrad_compute/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_defer_embedding_wgrad_compute/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_defer_embedding_wgrad_compute_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_defer_embedding_wgrad_compute_1node/model_config.yaml index f600f587ae9..340f6125883 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_defer_embedding_wgrad_compute_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_defer_embedding_wgrad_compute_1node/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_no_create_attention_mask_in_dataloader/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_no_create_attention_mask_in_dataloader/model_config.yaml index 58d9e2efbd2..2513ca33435 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_no_create_attention_mask_in_dataloader/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_no_create_attention_mask_in_dataloader/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_no_mmap_bin_files/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_no_mmap_bin_files/model_config.yaml index f1f26115d6b..9f563090212 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_no_mmap_bin_files/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_no_mmap_bin_files/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_reshard_1x4xNone/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_reshard_1x4xNone/model_config.yaml index ec29ea58ca6..bfb0c0f614d 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_reshard_1x4xNone/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_reshard_1x4xNone/model_config.yaml @@ -51,3 +51,4 @@ MODEL_ARGS: --use-persistent-ckpt-worker: true --verify-integrity: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_reshard_1x4xNone_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_reshard_1x4xNone_1node/model_config.yaml index a413b7ebb0c..346ef0e46dc 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_reshard_1x4xNone_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_pp2_resume_torch_dist_reshard_1x4xNone_1node/model_config.yaml @@ -50,3 +50,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_zp_z3_resume_fsdp_dtensor/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_zp_z3_resume_fsdp_dtensor/model_config.yaml index c07e943cbd8..3ef25d3f3af 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_zp_z3_resume_fsdp_dtensor/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_zp_z3_resume_fsdp_dtensor/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_zp_z3_resume_fsdp_dtensor_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_zp_z3_resume_fsdp_dtensor_1node/model_config.yaml index c07e943cbd8..3ef25d3f3af 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_zp_z3_resume_fsdp_dtensor_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp2_zp_z3_resume_fsdp_dtensor_1node/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce/model_config.yaml index d9cb8494444..a0c79e33ebb 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml index 11a72d75951..fe66ff34a7c 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce_param_gather_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce_param_gather_1node/model_config.yaml index 11a72d75951..fe66ff34a7c 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce_param_gather_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_dist_optimizer_overlap_grad_reduce_param_gather_1node/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_qk_layernorm_test_mode/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_qk_layernorm_test_mode/model_config.yaml index 9790a7f74ce..6b298d17c9e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_qk_layernorm_test_mode/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_qk_layernorm_test_mode/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce/model_config.yaml index bd6f526c209..f18736b62dc 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_1node/model_config.yaml index bd6f526c209..f18736b62dc 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_1node/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_qk_layernorm_test_mode/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_qk_layernorm_test_mode/model_config.yaml index addb335475f..bca8cea595b 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_qk_layernorm_test_mode/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_qk_layernorm_test_mode/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_qk_layernorm_test_mode_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_qk_layernorm_test_mode_1node/model_config.yaml index addb335475f..bca8cea595b 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_qk_layernorm_test_mode_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp1_resume_torch_dist_qk_layernorm_test_mode_1node/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_frozen_resume_torch_dist_reshard_8x1xNone/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_frozen_resume_torch_dist_reshard_8x1xNone/model_config.yaml index 2a4ce95cd6d..d76829ded35 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_frozen_resume_torch_dist_reshard_8x1xNone/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_frozen_resume_torch_dist_reshard_8x1xNone/model_config.yaml @@ -51,3 +51,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: frozen-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_resume_torch_dist_reshard_8x1xNone/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_resume_torch_dist_reshard_8x1xNone/model_config.yaml index 881fe7ebe2d..0a7e68912eb 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_resume_torch_dist_reshard_8x1xNone/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_resume_torch_dist_reshard_8x1xNone/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_resume_torch_dist_reshard_8x1xNone_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_resume_torch_dist_reshard_8x1xNone_1node/model_config.yaml index 630da63ed62..d1cc6d40025 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_resume_torch_dist_reshard_8x1xNone_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_te_tp4_pp2_resume_torch_dist_reshard_8x1xNone_1node/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml index f28a2d05a5c..c2da6641c92 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_fsdp2_resume_torch_dist_te/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_fsdp2_resume_torch_dist_te/model_config.yaml index c9eea9f5b0c..f0818de4b4e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_fsdp2_resume_torch_dist_te/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_fsdp2_resume_torch_dist_te/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml index 6588160cc67..7a103b83328 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp1_resume_torch_dist_dist_optimizer_overlap_grad_reduce_param_gather/model_config.yaml @@ -57,3 +57,4 @@ MODEL_ARGS: --use-persistent-ckpt-worker: true --verify-integrity: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2/model_config.yaml index 7ee44b85c81..07dc4d8e44b 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2_fp16/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2_fp16/model_config.yaml index 8e68343c17e..0763212e0e5 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2_fp16/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2_fp16/model_config.yaml @@ -51,3 +51,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2_resume_torch_dist/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2_resume_torch_dist/model_config.yaml index abe46384819..85ac5ccf314 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2_resume_torch_dist/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp2_resume_torch_dist/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp4/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp4/model_config.yaml index 57db98fd87e..57afdb9c140 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp4/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp4/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp4_resume_torch_dist/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp4_resume_torch_dist/model_config.yaml index 9bd0e50311d..feb482e71c8 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp4_resume_torch_dist/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp1_pp4_resume_torch_dist/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_resume_torch_dist_uninstall_te/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_resume_torch_dist_uninstall_te/model_config.yaml index a13ee609ace..661f2f722f5 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_resume_torch_dist_uninstall_te/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_resume_torch_dist_uninstall_te/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_uninstall_te/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_uninstall_te/model_config.yaml index d44645ba5e3..735b0bf3a18 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_uninstall_te/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_uninstall_te/model_config.yaml @@ -57,3 +57,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_uninstall_te_1node/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_uninstall_te_1node/model_config.yaml index d44645ba5e3..735b0bf3a18 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_uninstall_te_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp2_pp2_uninstall_te_1node/model_config.yaml @@ -57,3 +57,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1/model_config.yaml index abc889ac89e..3fc4281a27e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1/model_config.yaml @@ -51,3 +51,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1_resume_torch/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1_resume_torch/model_config.yaml index e61ad0b9ea9..0c2a0d64480 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1_resume_torch/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1_resume_torch/model_config.yaml @@ -48,3 +48,4 @@ MODEL_ARGS: --apply-query-key-layer-scaling: true --log-memory-to-tensorboard: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1_resume_torch_dist/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1_resume_torch_dist/model_config.yaml index a46065b1121..aff7990d847 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1_resume_torch_dist/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_mcore_tp4_pp1_resume_torch_dist/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_weekly_mcore_tp2_pp2_current_scaling_native_fp8_tp_pp_sp_tp_overlap/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_weekly_mcore_tp2_pp2_current_scaling_native_fp8_tp_pp_sp_tp_overlap/model_config.yaml index e8673fbae20..7f073c4d995 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_weekly_mcore_tp2_pp2_current_scaling_native_fp8_tp_pp_sp_tp_overlap/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_weekly_mcore_tp2_pp2_current_scaling_native_fp8_tp_pp_sp_tp_overlap/model_config.yaml @@ -63,3 +63,4 @@ METRICS: - iteration-time - lm loss - "mem-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt3_weekly_mcore_tp4_cp2_current_scaling_native_fp8_tp_sp_cp_tp_overlap/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt3_weekly_mcore_tp4_cp2_current_scaling_native_fp8_tp_sp_cp_tp_overlap/model_config.yaml index 2dccf144291..2772714757c 100644 --- a/tests/functional_tests/test_cases/gpt/gpt3_weekly_mcore_tp4_cp2_current_scaling_native_fp8_tp_sp_cp_tp_overlap/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt3_weekly_mcore_tp4_cp2_current_scaling_native_fp8_tp_sp_cp_tp_overlap/model_config.yaml @@ -57,3 +57,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_async_sched/golden_values_dev_dgx_gb200.json b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_async_sched/golden_values_dev_dgx_gb200.json new file mode 100644 index 00000000000..5737d10fce9 --- /dev/null +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_async_sched/golden_values_dev_dgx_gb200.json @@ -0,0 +1,711 @@ +{ + "0": { + "input_prompt": "The $500 Cup of coffee?\nConsider this, most Americans spend an average of $1,500-2,000 a year on this bean water.\nI have a few question for you: \nHow has business been the past few months?\nDo you ever feel like your business is stuck?\nDon't feel like you're able to improve performance and make changes required to achieve success ?\nAre your customers spneding less and less and wanting more?\nHave the gas prices affected your business?\nDo you have employees and do they hate you or wish they could quit?\n\nNow, before you and I can decide wheter or not I will be a good fit for your business we should talk this over with coffee.\nAnd, just to warn you this isn't some casual thing. This is not a date or time to be personal or social (but by all means share what you will coz I'll gladly listen).\nTher eare two major talking points and stratagies we will focios on in our lil coffee social\nFor one, we will find your unique selling Proposition (USP).\nDo have the best price? Are you the cheapest in town? Are your customers jerks? Do you haVE REGULARS? Why do people come back?\nwe'll also look for the holes in your business bucket. I'm willing to bet there's a hole or two in your business we'll find together that'll make this 500 dollar cup of Joe pay for itse;f immedietly.\nMany find themselves to be more profitable by just finding out where the dollars are escaping in their business and I like to think of myself as a guy that comes along with some spakel or putty and patch those holes up for you.\nBeleive me, just fixing one hole can mean a lot...just think about a sinking boat that has a hole in it that's about 3\u201d in diameter... it doesn't take long to sink.\nI have no agenda, besides f=getting to know your business and seeing wher I can patch the holes and find what makes you do darn unique (I know this won't take long.)\nMany folks, I bet, will find what they need to get off their chest with a quick phone call and they just paypal me the money and make a coffee at home. Look, that's fine too.\nI just to get you ot of your comfort zone, because this is where it all starts my frind.\nSome smart GOAT entrepreneur will probably get everything they need just out of our lil mini consulatant for the more extensive business consukting I offer, and look, that's fine, too.\nMaybe this cup of coffee will be all you need to gtet back on your feet, but not only surive, but thrive!\nI'm not trying to make frineds, or make a bunch of money, or look down your shirt\nBut this is only going to be a 45 minute (max) coffee chat\nAnd, I know you (and me) both have a lot on our plates. So no messing around\nAfter our consultation I will follow up with you in a few days and see how things are going, then I will be emailing you about once or twice every two weeks, just to make sure you're staying on task and implementing what we went over.\nTghere is no obligation to go any further and will gladly give you back your money if this pre-consultation doesn't put you on the right path or you don't get any value out of it...", + "generated_text": " $ is a$ is a $ is a $ is a $ is a $ is a $$1, you\n$ $$ $\n$ $- $\n$\n$\n$ $$$\n$\n$\n$\n$\n$\n$\n$\n$???????, $\n-You\n$\n$-1\n$\n$\nHow\n$.?\n\n$\n$\n1.1\nI\n$.\nI\n\n\nHow\n$,,,0,1,0,00- $\n... is\nThe\nThe\n", + "generated_tokens": [ + 1659, + 1395, + 1261, + 1036, + 1395, + 1261, + 1659, + 1395, + 1261, + 1659, + 1395, + 1261, + 1659, + 1395, + 1261, + 1659, + 1395, + 1261, + 1659, + 1036, + 1049, + 1044, + 1636, + 1010, + 1036, + 1659, + 1036, + 1659, + 1010, + 1036, + 1659, + 1045, + 1659, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1659, + 1036, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1063, + 1063, + 1063, + 1063, + 1063, + 1063, + 1063, + 1044, + 1659, + 1010, + 1045, + 4568, + 1010, + 1036, + 1010, + 1036, + 1045, + 1049, + 1010, + 1036, + 1010, + 1036, + 1010, + 7801, + 1010, + 1036, + 1046, + 9551, + 1010, + 1036, + 1010, + 1036, + 1010, + 1049, + 1046, + 1049, + 1010, + 1073, + 1010, + 1036, + 1046, + 1010, + 1073, + 1010, + 1010, + 1010, + 7801, + 1010, + 1036, + 1044, + 1044, + 1044, + 1048, + 1044, + 1049, + 1044, + 1048, + 1044, + 1048, + 1048, + 1045, + 1659, + 1010, + 2880, + 1395, + 1010, + 1784, + 1010, + 1784, + 1010 + ], + "latency": 15.173143148422241, + "ttft": 0.27222132682800293, + "cuda_graph_request_count_map": null, + "step_count": 132, + "top_n_logprobs": null, + "prompt_top_n_logprobs": null + }, + "32": { + "input_prompt": "create a conversational article", + "generated_text": " about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article.", + "generated_tokens": [ + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046 + ], + "latency": 15.158410787582397, + "ttft": 0.2701425552368164, + "cuda_graph_request_count_map": null, + "step_count": 132, + "top_n_logprobs": null, + "prompt_top_n_logprobs": null + }, + "64": { + "input_prompt": "App Concept: \"Eggy's Interactive Adventure World\"\n\nEggy's Interactive Adventure World is a cutting-edge app designed for phones and tablets that combines immersive storytelling with engaging, interactive experiences for both parents and children. This innovative app allows users to explore the world of Eggy and participate in a series of adventures that teach essential life skills, foster personal growth, and strengthen parent-child bonds.\n\nKey Features:\n\n1. Interactive Stories: The app brings the 20 Eggy stories to life through interactive, animated storytelling. Users can read, listen, and actively participate in the stories, making choices that guide Eggy through various adventures, each with unique lessons and experiences.\n2. Augmented Reality (AR) Integration: The app utilizes AR technology, allowing users to interact with Eggy and the story characters in their real-world environment. This immersive experience creates a deeper connection with the stories and encourages engagement.\n3. Personalized Adventures: The app customizes each story based on the user's interests, needs, and personal growth areas. By collecting user input, the app tailors the storylines and learning outcomes to ensure that each adventure resonates with the child's unique experiences.\n4. Parent-Child Cooperative Play: The app encourages parents and children to collaborate and engage in the stories together. Through a split-screen mode or multi-device connectivity, parents and children can actively participate in the stories and learn essential life skills together.\n5. Skill-Building Mini-Games: Interspersed throughout the stories are educational mini-games that reinforce the lessons and themes. These games are designed to be fun, engaging, and accessible to children, encouraging them to practice and develop various skills.\n6. Progress Tracking and Rewards: The app tracks the user's progress through the stories and rewards them with badges, trophies, and other collectibles. This gamification aspect encourages continued engagement and fosters a sense of accomplishment and growth.\n7. Community Features: The app provides a safe, moderated online community where parents can share their experiences, discuss the stories and lessons, and connect with other families who are also using the app. This community fosters support, inspiration, and learning.\n8. Expert Resources: The app includes access to resources from experts in child development, psychology, and education, providing parents with additional tools and guidance on fostering the growth and well-being of their children.\n\nEggy's Interactive Adventure World is a one-of-a-kind, mind-blowing app that revolutionizes the way parents and children engage with stories and learn essential life skills. Combining the power of storytelling, cutting-edge technology, and expert resources, this app is a must-have for families looking to strengthen their bonds and support their children's personal growth.", + "generated_text": "The 1999\u20132000 season was the 10th season of the National Hockey League (NHL). The 1999\u20132000 season was the 10th season of the National Hockey League (NHL). The 1999\u20132000 season was the 10th season of the National Hockey League (NHL). The 1999\u20132000 season was the 10th season was the 10th season was the 10 season was the 10 season was the 10", + "generated_tokens": [ + 2, + 1784, + 1032, + 1049, + 1057, + 1057, + 1057, + 1882, + 1050, + 1048, + 1048, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 1411, + 5526, + 1307, + 1278, + 5805, + 46982, + 8525, + 1319, + 19644, + 1076, + 1577, + 1531, + 1032, + 1049, + 1057, + 1057, + 1057, + 1882, + 1050, + 1048, + 1048, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 1411, + 5526, + 1307, + 1278, + 5805, + 46982, + 8525, + 1319, + 19644, + 1076, + 1577, + 1531, + 1032, + 1049, + 1057, + 1057, + 1057, + 1882, + 1050, + 1048, + 1048, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 1411, + 5526, + 1307, + 1278, + 5805, + 46982, + 8525, + 1319, + 19644, + 1076, + 1577, + 1531, + 1032, + 1049, + 1057, + 1057, + 1057, + 1882, + 1050, + 1048, + 1048, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 1411, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 1411, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048 + ], + "latency": 15.145662069320679, + "ttft": 0.26699352264404297, + "cuda_graph_request_count_map": null, + "step_count": 132, + "top_n_logprobs": null, + "prompt_top_n_logprobs": null + }, + "96": { + "input_prompt": "what is the concept of double materiality in sustainability?", + "generated_text": " What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double material", + "generated_tokens": [ + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318 + ], + "latency": 15.136452198028564, + "ttft": 0.24875521659851074, + "cuda_graph_request_count_map": null, + "step_count": 132, + "top_n_logprobs": null, + "prompt_top_n_logprobs": null + }, + "127": { + "input_prompt": "as before we talk about using , i usually do if else() to in case the property would be null before.\nif today i public a property and i want to use using to instead , will code be like this?\n\npublic Material Mat;\n\nusing(Mat){\n// Do something;\n}\n\nis that right?", + "generated_text": " i am not sure if i am using the wrong syntax or if i am using the syntax correctly but i am not sure if i am using the syntax correctly or if i am using the syntax incorrectly.\n\nI am not sure if i am using the wrong syntax or if i am using the syntax correctly but i am not sure if i am using the syntax correctly or if i am using the syntax incorrectly.\n\nI am not sure if i am using the wrong syntax or if i am using the syntax correctly but i am not sure if i am using the syntax correctly or if i am using the syntax incorrectly.\n\nI am not sure if i am using", + "generated_tokens": [ + 1623, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 8462, + 22692, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1809, + 1623, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 73751, + 1338, + 1073, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 8462, + 22692, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1809, + 1623, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 73751, + 1338, + 1073, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 8462, + 22692, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1809, + 1623, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 73751, + 1338, + 1073, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505 + ], + "latency": 15.125433683395386, + "ttft": 0.23802447319030762, + "cuda_graph_request_count_map": null, + "step_count": 132, + "top_n_logprobs": null, + "prompt_top_n_logprobs": null + }, + "throughput": [ + 836.2814759906043, + 1044.265060439428, + 1053.2105377765834 + ], + "mem-max-allocated-bytes": 22954507776, + "lifetime_prefill_token_count": 28887, + "async_sched_step_count": 131, + "async_sched_compaction_step_count": 4 +} \ No newline at end of file diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_async_sched/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_async_sched/golden_values_dev_dgx_h100.json new file mode 100644 index 00000000000..af2242f77a9 --- /dev/null +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_async_sched/golden_values_dev_dgx_h100.json @@ -0,0 +1,711 @@ +{ + "0": { + "input_prompt": "The $500 Cup of coffee?\nConsider this, most Americans spend an average of $1,500-2,000 a year on this bean water.\nI have a few question for you: \nHow has business been the past few months?\nDo you ever feel like your business is stuck?\nDon't feel like you're able to improve performance and make changes required to achieve success ?\nAre your customers spneding less and less and wanting more?\nHave the gas prices affected your business?\nDo you have employees and do they hate you or wish they could quit?\n\nNow, before you and I can decide wheter or not I will be a good fit for your business we should talk this over with coffee.\nAnd, just to warn you this isn't some casual thing. This is not a date or time to be personal or social (but by all means share what you will coz I'll gladly listen).\nTher eare two major talking points and stratagies we will focios on in our lil coffee social\nFor one, we will find your unique selling Proposition (USP).\nDo have the best price? Are you the cheapest in town? Are your customers jerks? Do you haVE REGULARS? Why do people come back?\nwe'll also look for the holes in your business bucket. I'm willing to bet there's a hole or two in your business we'll find together that'll make this 500 dollar cup of Joe pay for itse;f immedietly.\nMany find themselves to be more profitable by just finding out where the dollars are escaping in their business and I like to think of myself as a guy that comes along with some spakel or putty and patch those holes up for you.\nBeleive me, just fixing one hole can mean a lot...just think about a sinking boat that has a hole in it that's about 3\u201d in diameter... it doesn't take long to sink.\nI have no agenda, besides f=getting to know your business and seeing wher I can patch the holes and find what makes you do darn unique (I know this won't take long.)\nMany folks, I bet, will find what they need to get off their chest with a quick phone call and they just paypal me the money and make a coffee at home. Look, that's fine too.\nI just to get you ot of your comfort zone, because this is where it all starts my frind.\nSome smart GOAT entrepreneur will probably get everything they need just out of our lil mini consulatant for the more extensive business consukting I offer, and look, that's fine, too.\nMaybe this cup of coffee will be all you need to gtet back on your feet, but not only surive, but thrive!\nI'm not trying to make frineds, or make a bunch of money, or look down your shirt\nBut this is only going to be a 45 minute (max) coffee chat\nAnd, I know you (and me) both have a lot on our plates. So no messing around\nAfter our consultation I will follow up with you in a few days and see how things are going, then I will be emailing you about once or twice every two weeks, just to make sure you're staying on task and implementing what we went over.\nTghere is no obligation to go any further and will gladly give you back your money if this pre-consultation doesn't put you on the right path or you don't get any value out of it...", + "generated_text": " $ is a$ is a $ is a $ is a $ is a $ is a $$1, you\n$ $$ $\n$ $- $\n$\n$\n$ $$$\n$\n$\n$\n$\n$\n$\n$\n$???????, $\n-You\n$\n$-1\n$\n$\nHow\n$.?\n\n$\n$\n1.1\nI\n$.\nI\n\n\nHow\n$,,,0,1,0,00- $\n... is\nThe\nThe\n", + "generated_tokens": [ + 1659, + 1395, + 1261, + 1036, + 1395, + 1261, + 1659, + 1395, + 1261, + 1659, + 1395, + 1261, + 1659, + 1395, + 1261, + 1659, + 1395, + 1261, + 1659, + 1036, + 1049, + 1044, + 1636, + 1010, + 1036, + 1659, + 1036, + 1659, + 1010, + 1036, + 1659, + 1045, + 1659, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1659, + 1036, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1010, + 1036, + 1063, + 1063, + 1063, + 1063, + 1063, + 1063, + 1063, + 1044, + 1659, + 1010, + 1045, + 4568, + 1010, + 1036, + 1010, + 1036, + 1045, + 1049, + 1010, + 1036, + 1010, + 1036, + 1010, + 7801, + 1010, + 1036, + 1046, + 9551, + 1010, + 1036, + 1010, + 1036, + 1010, + 1049, + 1046, + 1049, + 1010, + 1073, + 1010, + 1036, + 1046, + 1010, + 1073, + 1010, + 1010, + 1010, + 7801, + 1010, + 1036, + 1044, + 1044, + 1044, + 1048, + 1044, + 1049, + 1044, + 1048, + 1044, + 1048, + 1048, + 1045, + 1659, + 1010, + 2880, + 1395, + 1010, + 1784, + 1010, + 1784, + 1010 + ], + "latency": 10.892494678497314, + "ttft": 0.194596529006958, + "cuda_graph_request_count_map": null, + "step_count": 132, + "top_n_logprobs": null, + "prompt_top_n_logprobs": null + }, + "32": { + "input_prompt": "create a conversational article", + "generated_text": " about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article. The article should be about the topic of the article.", + "generated_tokens": [ + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046, + 1531, + 9369, + 2715, + 1402, + 2314, + 1278, + 17915, + 1307, + 1278, + 9369, + 1046 + ], + "latency": 10.882038831710815, + "ttft": 0.1947627067565918, + "cuda_graph_request_count_map": null, + "step_count": 132, + "top_n_logprobs": null, + "prompt_top_n_logprobs": null + }, + "64": { + "input_prompt": "App Concept: \"Eggy's Interactive Adventure World\"\n\nEggy's Interactive Adventure World is a cutting-edge app designed for phones and tablets that combines immersive storytelling with engaging, interactive experiences for both parents and children. This innovative app allows users to explore the world of Eggy and participate in a series of adventures that teach essential life skills, foster personal growth, and strengthen parent-child bonds.\n\nKey Features:\n\n1. Interactive Stories: The app brings the 20 Eggy stories to life through interactive, animated storytelling. Users can read, listen, and actively participate in the stories, making choices that guide Eggy through various adventures, each with unique lessons and experiences.\n2. Augmented Reality (AR) Integration: The app utilizes AR technology, allowing users to interact with Eggy and the story characters in their real-world environment. This immersive experience creates a deeper connection with the stories and encourages engagement.\n3. Personalized Adventures: The app customizes each story based on the user's interests, needs, and personal growth areas. By collecting user input, the app tailors the storylines and learning outcomes to ensure that each adventure resonates with the child's unique experiences.\n4. Parent-Child Cooperative Play: The app encourages parents and children to collaborate and engage in the stories together. Through a split-screen mode or multi-device connectivity, parents and children can actively participate in the stories and learn essential life skills together.\n5. Skill-Building Mini-Games: Interspersed throughout the stories are educational mini-games that reinforce the lessons and themes. These games are designed to be fun, engaging, and accessible to children, encouraging them to practice and develop various skills.\n6. Progress Tracking and Rewards: The app tracks the user's progress through the stories and rewards them with badges, trophies, and other collectibles. This gamification aspect encourages continued engagement and fosters a sense of accomplishment and growth.\n7. Community Features: The app provides a safe, moderated online community where parents can share their experiences, discuss the stories and lessons, and connect with other families who are also using the app. This community fosters support, inspiration, and learning.\n8. Expert Resources: The app includes access to resources from experts in child development, psychology, and education, providing parents with additional tools and guidance on fostering the growth and well-being of their children.\n\nEggy's Interactive Adventure World is a one-of-a-kind, mind-blowing app that revolutionizes the way parents and children engage with stories and learn essential life skills. Combining the power of storytelling, cutting-edge technology, and expert resources, this app is a must-have for families looking to strengthen their bonds and support their children's personal growth.", + "generated_text": "The 1999\u20132000 season was the 10th season of the National Hockey League (NHL). The 1999\u20132000 season was the 10th season of the National Hockey League (NHL). The 1999\u20132000 season was the 10th season of the National Hockey League (NHL). The 1999\u20132000 season was the 10th season was the 10th season was the 10 season was the 10 season was the 10", + "generated_tokens": [ + 2, + 1784, + 1032, + 1049, + 1057, + 1057, + 1057, + 1882, + 1050, + 1048, + 1048, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 1411, + 5526, + 1307, + 1278, + 5805, + 46982, + 8525, + 1319, + 19644, + 1076, + 1577, + 1531, + 1032, + 1049, + 1057, + 1057, + 1057, + 1882, + 1050, + 1048, + 1048, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 1411, + 5526, + 1307, + 1278, + 5805, + 46982, + 8525, + 1319, + 19644, + 1076, + 1577, + 1531, + 1032, + 1049, + 1057, + 1057, + 1057, + 1882, + 1050, + 1048, + 1048, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 1411, + 5526, + 1307, + 1278, + 5805, + 46982, + 8525, + 1319, + 19644, + 1076, + 1577, + 1531, + 1032, + 1049, + 1057, + 1057, + 1057, + 1882, + 1050, + 1048, + 1048, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 1411, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 1411, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048, + 5526, + 1486, + 1278, + 1032, + 1049, + 1048 + ], + "latency": 10.872116804122925, + "ttft": 0.19449090957641602, + "cuda_graph_request_count_map": null, + "step_count": 132, + "top_n_logprobs": null, + "prompt_top_n_logprobs": null + }, + "96": { + "input_prompt": "what is the concept of double materiality in sustainability?", + "generated_text": " What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double materiality in sustainability? What is the concept of double material", + "generated_tokens": [ + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318, + 1532, + 1294, + 60187, + 1063, + 5675, + 1395, + 1278, + 7401, + 1307, + 5558, + 4318 + ], + "latency": 10.860994100570679, + "ttft": 0.18096256256103516, + "cuda_graph_request_count_map": null, + "step_count": 132, + "top_n_logprobs": null, + "prompt_top_n_logprobs": null + }, + "127": { + "input_prompt": "as before we talk about using , i usually do if else() to in case the property would be null before.\nif today i public a property and i want to use using to instead , will code be like this?\n\npublic Material Mat;\n\nusing(Mat){\n// Do something;\n}\n\nis that right?", + "generated_text": " i am not sure if i am using the wrong syntax or if i am using the syntax correctly but i am not sure if i am using the syntax correctly or if i am using the syntax incorrectly.\n\nI am not sure if i am using the wrong syntax or if i am using the syntax correctly but i am not sure if i am using the syntax correctly or if i am using the syntax incorrectly.\n\nI am not sure if i am using the wrong syntax or if i am using the syntax correctly but i am not sure if i am using the syntax correctly or if i am using the syntax incorrectly.\n\nI am not sure if i am using", + "generated_tokens": [ + 1623, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 8462, + 22692, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1809, + 1623, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 73751, + 1338, + 1073, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 8462, + 22692, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1809, + 1623, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 73751, + 1338, + 1073, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 8462, + 22692, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1809, + 1623, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 17047, + 1505, + 1693, + 1623, + 1855, + 2505, + 1278, + 22692, + 73751, + 1338, + 1073, + 1855, + 1605, + 5257, + 1693, + 1623, + 1855, + 2505 + ], + "latency": 10.852606058120728, + "ttft": 0.17280030250549316, + "cuda_graph_request_count_map": null, + "step_count": 132, + "top_n_logprobs": null, + "prompt_top_n_logprobs": null + }, + "throughput": [ + 1130.0929743054994, + 1467.7527907005106, + 1467.8709552178896 + ], + "mem-max-allocated-bytes": 22954507776, + "lifetime_prefill_token_count": 28887, + "async_sched_step_count": 131, + "async_sched_compaction_step_count": 4 +} \ No newline at end of file diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_async_sched/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_async_sched/model_config.yaml new file mode 100644 index 00000000000..e4c49d32c68 --- /dev/null +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_async_sched/model_config.yaml @@ -0,0 +1,59 @@ +ENV_VARS: + CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 + NCCL_ALGO: Ring + CUBLAS_WORKSPACE_CONFIG: :4096:8 +TEST_TYPE: frozen-start +MODE: inference +MODEL_ARGS: + --tiktoken-pattern: v2 + --use-mcore-models: true + --tokenizer-type: TikTokenizer + --tokenizer-model: ${CHECKPOINT_LOAD_PATH}/model/mcore_mistral/nemo_minitron-0.5b/v1/multiMixV8.gpt4o_nc_sd.500000.128k.vocab.json + --auto-detect-ckpt-format: true + --max-tokens-to-oom: 3600000 + --inference-max-seq-length: 4096 + --attention-backend: flash + --use-checkpoint-args: true + --micro-batch-size: 1 + --no-load-optim: true + --no-use-tokenizer-model-from-checkpoint-args: true + --timing-log-level: 0 + --load: ${CHECKPOINT_LOAD_PATH}/model/mcore_mistral/nemo_minitron-0.5b/v1 + --distributed-backend: nccl + --log-interval: 1 + --transformer-impl: transformer_engine + --tensor-model-parallel-size: 1 + --pipeline-model-parallel-size: 1 + --deterministic-mode: true + --ckpt-format: torch_dist + --bf16: true + --log-memory-to-tensorboard: true + --log-num-zeros-in-grad: true + --log-validation-ppl-to-tensorboard: true + --log-timers-to-tensorboard: true + --num-layers: 24 + --hidden-size: 1152 + --num-attention-heads: 16 + --max-position-embeddings: 1024 + --seq-length: 1024 + --temperature: 1.0 + --top_k: 1 + # Async scheduling only supports greedy sampling (top_k=1, top_p=0.0) and does + # not support log probabilities, stop words, chunked prefill, or prefix + # caching (see dynamic_engine._validate_async_sched_support_for_request). + --inference-dynamic-batching-buffer-size-gb: 20 + --inference-dynamic-batching-async-sched-mode: async + --dist-ckpt-strictness: log_unexpected + --inference-ckpt-non-strict: true # To handle the extra_state errors + --output-path: ${INFERENCE_OUTPUT_PATH} + --output-every-n-results: 32 + --prompt-file: ${DATA_PATH}/text/sharegpt-vicuna/filtered/processed.jsonl + --prompt-file-num-truncate: 128 # originally 1024 + --num-tokens-to-generate: 128 # originally 512 + --incoming-requests-per-step: 32 + --termination-id: -1 + --inference-repeat-n: 3 + --inference-logging-step-interval: 1 +METRICS: + - "generated_tokens" diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_chunked_prefill/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_chunked_prefill/model_config.yaml index c304e8bf5df..98215b5c3c4 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_chunked_prefill/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_chunked_prefill/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_chunked_prefill_cuda_graphs/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_chunked_prefill_cuda_graphs/model_config.yaml index 4b9e265c022..5046b7088e8 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_chunked_prefill_cuda_graphs/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_chunked_prefill_cuda_graphs/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_fp8_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_fp8_logitsmatch/model_config.yaml index 21fd7749ea7..3c00126bb32 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_fp8_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_fp8_logitsmatch/model_config.yaml @@ -30,7 +30,6 @@ MODEL_ARGS: --bf16: true --fp8-recipe: tensorwise --fp8-format: hybrid - --fp8-param-gather: true --first-last-layers-bf16: true --log-memory-to-tensorboard: true --log-num-zeros-in-grad: true @@ -58,3 +57,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_logitsmatch_decode_graphs_only/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_logitsmatch_decode_graphs_only/model_config.yaml index 9996433cf1f..d44ea0486e9 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_logitsmatch_decode_graphs_only/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_logitsmatch_decode_graphs_only/model_config.yaml @@ -30,7 +30,6 @@ MODEL_ARGS: --bf16: true --fp8-recipe: tensorwise --fp8-format: hybrid - --fp8-param-gather: true --first-last-layers-bf16: true --log-memory-to-tensorboard: true --log-num-zeros-in-grad: true @@ -59,3 +58,4 @@ METRICS: - "generated_tokens" - "logprobs" - "throughput" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_validation/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_validation/model_config.yaml index 2ac5db11472..fbafc91b4f8 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_validation/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_cuda_graphs_validation/model_config.yaml @@ -5,3 +5,4 @@ ENV_VARS: CUBLAS_WORKSPACE_CONFIG: :4096:8 MODEL_ARGS: TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_flashinfer/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_flashinfer/model_config.yaml index 90e1cf11107..0e8bfa319ad 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_flashinfer/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_flashinfer/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_logitsmatch/model_config.yaml index 4e250b956ef..e69e3e93703 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_logitsmatch/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching/model_config.yaml index 5cf5f3c9902..7a4a8902984 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill/model_config.yaml index cc4f4cbdd7c..30d9b109c09 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill_cuda_graphs/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill_cuda_graphs/model_config.yaml index 0a041e16316..3d3f3e5d372 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill_cuda_graphs/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill_cuda_graphs/model_config.yaml @@ -57,3 +57,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill_flashinfer/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill_flashinfer/model_config.yaml index c6d8a77c9be..d0fb2e0c1a2 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill_flashinfer/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_chunked_prefill_flashinfer/model_config.yaml @@ -57,3 +57,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_cuda_graphs/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_cuda_graphs/model_config.yaml index bd46560d562..cafcf760c76 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_cuda_graphs/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_cuda_graphs/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_lru/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_lru/model_config.yaml index efc71f68f70..a64b50f4d48 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_lru/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_prefix_caching_lru/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_stop_words/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_stop_words/model_config.yaml index 8259814da8b..1f117df5d63 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_stop_words/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_stop_words/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_top_n_logprobs/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_top_n_logprobs/model_config.yaml index 487ddc157f1..64989c12adf 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_top_n_logprobs/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_top_n_logprobs/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_top_p_sampling/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_top_p_sampling/model_config.yaml index 3e92838c537..b1204cfcca0 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_top_p_sampling/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_top_p_sampling/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_uvm_level1/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_uvm_level1/model_config.yaml index a5d304eff91..4b64cface52 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_uvm_level1/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_583m_uvm_level1/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml index d84dd24487f..dea5eec83dc 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_round_robin_zmq/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_load_balanced_zmq/golden_values_dev_dgx_h100.json similarity index 100% rename from tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_round_robin_zmq/golden_values_dev_dgx_h100.json rename to tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_load_balanced_zmq/golden_values_dev_dgx_h100.json diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_round_robin_zmq/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_load_balanced_zmq/model_config.yaml similarity index 98% rename from tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_round_robin_zmq/model_config.yaml rename to tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_load_balanced_zmq/model_config.yaml index 9ba47a56e6e..452dce08eae 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_round_robin_zmq/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_load_balanced_zmq/model_config.yaml @@ -44,7 +44,7 @@ MODEL_ARGS: --num-tokens-to-generate: 30 --inference-dynamic-batching-buffer-size-gb: 20 --inference-dynamic-batching-prefix-caching: true - --inference-dynamic-batching-prefix-caching-coordinator-policy: round_robin + --inference-dynamic-batching-prefix-caching-coordinator-policy: load_balanced --dist-ckpt-strictness: log_unexpected --inference-ckpt-non-strict: true # To handle the extra_state errors --output-path: ${INFERENCE_OUTPUT_PATH} @@ -55,3 +55,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_longest_prefix_zmq/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_longest_prefix_zmq/model_config.yaml index fcfc6c716f0..f6e3112b8dd 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_longest_prefix_zmq/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_longest_prefix_zmq/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp8_583m_prefix_caching/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp8_583m_prefix_caching/model_config.yaml index 72483d72ccb..23969d560c4 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp8_583m_prefix_caching/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp8_583m_prefix_caching/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp8_dp1_583m_logitsmatch_zmq/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp8_dp1_583m_logitsmatch_zmq/model_config.yaml index 345fc250694..3fe997dd46b 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp8_dp1_583m_logitsmatch_zmq/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp1_pp8_dp1_583m_logitsmatch_zmq/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_chunked_prefill/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_chunked_prefill/model_config.yaml index 51c1a4ad703..98be4420000 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_chunked_prefill/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_chunked_prefill/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_cuda_graphs/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_cuda_graphs/model_config.yaml index 83f68909f47..6101d34c87a 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_cuda_graphs/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_cuda_graphs/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_prefix_caching/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_prefix_caching/model_config.yaml index a157e899c2f..8df34df686e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_prefix_caching/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_prefix_caching/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_prefix_caching_cuda_graphs/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_prefix_caching_cuda_graphs/model_config.yaml index e1a1f680c0f..507e8abc0bf 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_prefix_caching_cuda_graphs/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_583m_prefix_caching_cuda_graphs/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_dp2_583m_logitsmatch_zmq/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_dp2_583m_logitsmatch_zmq/model_config.yaml index 3b55b09e82e..a3976951850 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_dp2_583m_logitsmatch_zmq/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp2_pp2_dp2_583m_logitsmatch_zmq/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp4_pp1_583m_flashinfer/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp4_pp1_583m_flashinfer/model_config.yaml index ea1d201f339..3791f475136 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp4_pp1_583m_flashinfer/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp4_pp1_583m_flashinfer/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_583m_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_583m_logitsmatch/model_config.yaml index 4458edf5772..451b5070661 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_583m_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_583m_logitsmatch/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_583m_prefix_caching/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_583m_prefix_caching/model_config.yaml index 8f8b9d55925..b1f473124ab 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_583m_prefix_caching/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_583m_prefix_caching/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_dp1_583m_logitsmatch_zmq/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_dp1_583m_logitsmatch_zmq/model_config.yaml index 88a3e40a193..e34e63ec4c0 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_dp1_583m_logitsmatch_zmq/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_dynamic_inference_tp8_pp1_dp1_583m_logitsmatch_zmq/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_grpo_basic_function/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_grpo_basic_function/model_config.yaml index 0143a39f017..f052a0bca75 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_grpo_basic_function/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_grpo_basic_function/model_config.yaml @@ -97,3 +97,4 @@ MODEL_ARGS: --finetune: true --inference-logging-step-interval: 1 METRICS: [] +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_grpo_tp1tp2_pp1_dp8_583m_throughputtest/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_grpo_tp1tp2_pp1_dp8_583m_throughputtest/model_config.yaml index 4f9be214289..7ef22862c00 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_grpo_tp1tp2_pp1_dp8_583m_throughputtest/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_grpo_tp1tp2_pp1_dp8_583m_throughputtest/model_config.yaml @@ -80,3 +80,4 @@ MODEL_ARGS: METRICS: - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_grpo_tp1tp2_pp1_dp8_583m_throughputtest_github/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_grpo_tp1tp2_pp1_dp8_583m_throughputtest_github/model_config.yaml index c8fa19d0500..e12eea86625 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_grpo_tp1tp2_pp1_dp8_583m_throughputtest_github/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_grpo_tp1tp2_pp1_dp8_583m_throughputtest_github/model_config.yaml @@ -86,4 +86,5 @@ METRICS: - "mem-max-allocated-bytes" THROUGHPUT_TEST_PARAMS: - --start_step: 1 \ No newline at end of file + --start_step: 1 +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_cudagraphs_throughput/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_cudagraphs_throughput/model_config.yaml index 654df68947f..d1e0a562625 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_cudagraphs_throughput/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_cudagraphs_throughput/model_config.yaml @@ -113,3 +113,4 @@ METRICS: - "mem-allocated-bytes" - "mem-max-allocated-bytes" - "iteration-time" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_throughput/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_throughput/model_config.yaml index b7fb41046f3..35b56bbe8f9 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_throughput/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_throughput/model_config.yaml @@ -112,3 +112,4 @@ METRICS: - "mem-allocated-bytes" - "mem-max-allocated-bytes" - "iteration-time" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_throughput_github/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_throughput_github/model_config.yaml index cc25f3ab90e..8b4a14b9836 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_throughput_github/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_grpo_tp4_pp1_dp2_8b_throughput_github/model_config.yaml @@ -98,3 +98,4 @@ METRICS: - "mem-allocated-bytes" - "mem-max-allocated-bytes" - "iteration-time" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_offline_inference_async_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_offline_inference_async_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml index 076659eb455..2a91ac2a559 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_offline_inference_async_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_offline_inference_async_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml @@ -60,3 +60,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_offline_inference_sync_tp1_pp1_583m_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_offline_inference_sync_tp1_pp1_583m_logitsmatch/model_config.yaml index 6a5fc63aab3..bf82a3b32d9 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_offline_inference_sync_tp1_pp1_583m_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_offline_inference_sync_tp1_pp1_583m_logitsmatch/model_config.yaml @@ -60,3 +60,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_offline_inference_sync_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_offline_inference_sync_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml index ba3622a9417..5f18ceddcf9 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_offline_inference_sync_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_offline_inference_sync_tp1_pp1_dp8_583m_logitsmatch_zmq/model_config.yaml @@ -59,3 +59,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_16b_multiprompt_tokensmatch/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_16b_multiprompt_tokensmatch/model_config.yaml index 40b45024cb1..01680ec5e2e 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_16b_multiprompt_tokensmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_16b_multiprompt_tokensmatch/model_config.yaml @@ -1,5 +1,6 @@ ENV_VARS: CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE: 1 NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 NCCL_ALGO: Ring CUBLAS_WORKSPACE_CONFIG: :4096:8 @@ -83,3 +84,4 @@ MODEL_ARGS: --inference-dynamic-batching-buffer-size-gb: 20 METRICS: - "generated_text" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_cudagraphs/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_cudagraphs/model_config.yaml index 9a47281703a..28635a7fec7 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_cudagraphs/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_cudagraphs/model_config.yaml @@ -1,5 +1,6 @@ ENV_VARS: CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE: 1 NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 NCCL_ALGO: Ring CUBLAS_WORKSPACE_CONFIG: :4096:8 @@ -55,3 +56,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_fp8_cudagraphs/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_fp8_cudagraphs/model_config.yaml index 99bcc433ad1..6eaff84d57f 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_fp8_cudagraphs/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_fp8_cudagraphs/model_config.yaml @@ -1,5 +1,6 @@ ENV_VARS: CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE: 1 NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 NCCL_ALGO: Ring CUBLAS_WORKSPACE_CONFIG: :4096:8 @@ -60,3 +61,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_logitsmatch/model_config.yaml index 1c78b466b1e..bac8adc5051 100644 --- a/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/gpt/gpt_static_inference_tp1_pp1_583m_logitsmatch/model_config.yaml @@ -1,5 +1,6 @@ ENV_VARS: CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE: 1 NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 NCCL_ALGO: Ring CUBLAS_WORKSPACE_CONFIG: :4096:8 @@ -50,3 +51,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_ep8_nanov3_chunked_prefill/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_ep8_nanov3_chunked_prefill/model_config.yaml index ba07a85c024..9cce670aa6b 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_ep8_nanov3_chunked_prefill/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_ep8_nanov3_chunked_prefill/model_config.yaml @@ -66,3 +66,4 @@ MODEL_ARGS: --inference-moe-token-dispatcher-type: nvls --inference-logging-step-interval: 1 METRICS: +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_2b_async_sched_async/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_2b_async_sched_async/golden_values_dev_dgx_h100.json new file mode 100644 index 00000000000..38f7ec167de --- /dev/null +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_2b_async_sched_async/golden_values_dev_dgx_h100.json @@ -0,0 +1,104 @@ +{ + "0": { + "generated_tokens": [ + 2157, + 1395, + 1605, + 1046, + 2157, + 1395, + 1261, + 3535, + 2478, + 1636, + 1710, + 2012, + 1261, + 6854, + 1435, + 1261, + 1289, + 1490, + 3206, + 1044, + 1321, + 2478, + 1636, + 1710, + 2012, + 1261, + 6854, + 1435, + 1261, + 15353 + ] + }, + "1": { + "generated_tokens": [ + 2157, + 1395, + 1605, + 1046, + 2157, + 1395, + 1261, + 3535, + 2478, + 1636, + 1710, + 2012, + 1261, + 6854, + 1435, + 1261, + 1289, + 1490, + 3206, + 1044, + 1321, + 1636, + 1710, + 2012, + 1261, + 6854, + 1435, + 1261, + 1289, + 1490 + ] + }, + "2": { + "generated_tokens": [ + 2157, + 1395, + 1605, + 1046, + 2157, + 1395, + 1261, + 3535, + 2478, + 1636, + 1710, + 2012, + 1261, + 6854, + 1435, + 1261, + 1289, + 1490, + 3206, + 1044, + 1321, + 2478, + 1636, + 1710, + 2012, + 1261, + 6854, + 1435, + 1261, + 15353 + ] + } +} diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_2b_async_sched_async/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_2b_async_sched_async/model_config.yaml new file mode 100644 index 00000000000..3cacaf01867 --- /dev/null +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_2b_async_sched_async/model_config.yaml @@ -0,0 +1,75 @@ +ENV_VARS: + CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 + NCCL_ALGO: Ring + CUBLAS_WORKSPACE_CONFIG: :4096:8 + TRITON_CACHE_AUTOTUNING: 0 + MAMBA_DETERMINISTIC: 1 +TEST_TYPE: frozen-start +MODE: inference +MODEL_ARGS: + --log-num-zeros-in-grad: true + --log-validation-ppl-to-tensorboard: true + --log-timers-to-tensorboard: true + --log-memory-to-tensorboard: true + --timing-log-level: 0 + --load: ${CHECKPOINT_LOAD_PATH}/model/mamba_hybrid_2b/dcp/mcore-v1_bf16/checkpoint + --tokenizer-model: ${CHECKPOINT_LOAD_PATH}/model/mamba_hybrid_2b/dcp/mcore-v1_bf16/multiMixV8.gpt4o_nc_sd.500000.128k.vocab.json + --tokenizer-type: TikTokenizer + --tiktoken-pattern: v2 + --distributed-backend: nccl + --log-interval: 1 + --transformer-impl: transformer_engine + --tensor-model-parallel-size: 1 + --pipeline-model-parallel-size: 1 + --expert-model-parallel-size: 1 + --use-mcore-models: true + --model-provider: hybrid + --init-method-std: 0.0198 + --untie-embeddings-and-output-weights: true + --disable-bias-linear: true + --init-method-std: 0.014 + --position-embedding-type: none + --hidden-size: 2048 + --ffn-hidden-size: 11264 + --num-attention-heads: 16 + --kv-channels: 128 + --hybrid-layer-pattern: M-M-M-M*-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M*-M-M-M-M-M- + --spec: megatron.core.models.hybrid.hybrid_layer_specs hybrid_stack_spec + --normalization: RMSNorm + --swiglu: true + --attention-dropout: 0.0 + --hidden-dropout: 0.0 + --seq-length: 4096 + --max-position-embeddings: 4096 + --micro-batch-size: 1 + --ckpt-format: torch_dist + --ckpt-fully-parallel-save: true + --ckpt-fully-parallel-load: true + --ckpt-assume-constant-structure: true + --dist-ckpt-strictness: log_unexpected + --bf16: true + --attention-backend: flash + --no-create-attention-mask-in-dataloader: true + --num-workers: 8 + --use-checkpoint-args: true + --no-use-tokenizer-model-from-checkpoint-args: true + --no-load-optim: true + --deterministic-mode: true + --save-interval: 2000 + --temperature: 1.0 + --top_k: 1 + --num-tokens-to-generate: 30 + --max-tokens-to-oom: 3600000 + --inference-max-seq-length: 4096 + --output-path: ${INFERENCE_OUTPUT_PATH} + --prompt-file: ./tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_2b_async_sched_async/prompts.jsonl + --incoming-requests-per-step: 1 + --inference-repeat-n: 2 + --no-record-throughput: true + --mamba-inference-conv-states-dtype: fp32 + --mamba-inference-ssm-states-dtype: fp32 + --inference-dynamic-batching-async-sched-mode: async +METRICS: + - "generated_tokens" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_2b_async_sched_async/prompts.jsonl b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_2b_async_sched_async/prompts.jsonl new file mode 100644 index 00000000000..e5869299d3a --- /dev/null +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_2b_async_sched_async/prompts.jsonl @@ -0,0 +1,3 @@ +{"text":"Time travel to 2008, and go to a bar or a club or one of the myriad disco-basements on the Lower East Side that does not quite know which of those it is. Dance awkwardly in a room full of other glittered-up nerds, and wait for something to happen, buoyed on the feeling that this is the big swollen heart of life, that this is New York like the movies."} +{"text":"Time travel to 2008, and go to a bar or a club or one of the myriad disco-basements on the Lower East Side that does not quite know which of those it is. Dance awkwardly in a room full of other glittered-up nerds, and wait for something to happen, buoyed on the feeling that this is the big swollen heart of life, that this is New York like the movies."} +{"text":"Time travel to 2008, and go to a bar or a club or one of the myriad disco-basements on the Lower East Side that does not quite know which of those it is. Dance awkwardly in a room full of other glittered-up nerds, and wait for something to happen, buoyed on the feeling that this is the big swollen heart of life, that this is New York like the movies."} diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m/model_config.yaml index 4b258afe0d6..5bf8fc30d2a 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m/model_config.yaml @@ -73,3 +73,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_chunked_prefill/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_chunked_prefill/model_config.yaml index bd86d2faa44..5f4cd9cef83 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_chunked_prefill/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_chunked_prefill/model_config.yaml @@ -76,3 +76,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_flashinfer/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_flashinfer/golden_values_dev_dgx_h100.json index 956754f44c1..b79d1ed4115 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_flashinfer/golden_values_dev_dgx_h100.json +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_flashinfer/golden_values_dev_dgx_h100.json @@ -1,41 +1,41 @@ { "0": { "input_prompt": "Time travel to 2008, and go to a bar or a club or one of the myriad disco-basements on the Lower East Side that does not quite know which of those it is. Dance awkwardly in a room full of other glittered-up nerds, and wait for something to happen, buoyed on the feeling that this is the big swollen heart of life, that this is New York like the movies.", - "generated_text": " You are not alone. You are not alone. You are not alone. You are not alone. You are not alone. You are not alone.", + "generated_text": " You are a part of it, and you are not. You are a part of it, and you are not. You are a part of it", "generated_tokens": [ 3213, 1584, - 1605, - 9412, - 1046, - 3213, + 1261, + 1805, + 1307, + 1494, + 1044, + 1321, + 1636, 1584, 1605, - 9412, 1046, 3213, 1584, - 1605, - 9412, - 1046, - 3213, + 1261, + 1805, + 1307, + 1494, + 1044, + 1321, + 1636, 1584, 1605, - 9412, 1046, 3213, 1584, - 1605, - 9412, - 1046, - 3213, - 1584, - 1605, - 9412, - 1046 + 1261, + 1805, + 1307, + 1494 ], - "latency": 1.878054141998291, - "ttft": 0.07786321640014648, + "latency": 1.882845401763916, + "ttft": 0.0774838924407959, "cuda_graph_request_count_map": null, "step_count": 30, "top_n_logprobs": null, @@ -133,33 +133,33 @@ -3.0331737995147705, -1.9080564975738525, -2.52506947517395, - -2.325258493423462, - -1.180279016494751, - -1.1824196577072144, - -0.39788734912872314, - -1.110222578048706, - -1.5034958124160767, - -0.9765141606330872, - -0.9300433397293091, - -0.15196305513381958, - -0.2200203537940979, - -0.06051275506615639, - -0.6840062737464905, - -1.0964292287826538, - -0.17654964327812195, - -0.18547140061855316, - -0.06710249185562134, - -0.4758152365684509, - -0.6657928228378296, - -0.10342729091644287, - -0.10059614479541779, - -0.046978313475847244, - -0.410809725522995, - -0.428723007440567, - -0.06053968518972397, - -0.06518109142780304, - -0.030038274824619293, - -0.3271780014038086 + -3.010035514831543, + -0.011503085494041443, + -1.3661863803863525, + -0.7587734460830688, + -1.4640830755233765, + -1.4567692279815674, + -1.0451016426086426, + -2.3799993991851807, + -1.3697059154510498, + -1.0724036693572998, + -0.5331835150718689, + -1.8569976091384888, + -0.9731019735336304, + -0.05548504367470741, + -0.29222357273101807, + -0.2191229909658432, + -0.23294074833393097, + -0.6115235090255737, + -0.39632147550582886, + -1.1302311420440674, + -0.3844905197620392, + -0.7445144653320312, + -0.20952190458774567, + -0.29755353927612305, + -0.0599716454744339, + -0.008489826694130898, + -0.03623323515057564 ], "logprobs": [ -9.498085021972656, @@ -252,35 +252,37 @@ -3.0331737995147705, -1.9080564975738525, -2.52506947517395, - -2.325258493423462, - -1.180279016494751, - -1.1824196577072144, - -0.39788734912872314, - -1.110222578048706, - -1.5034958124160767, - -0.9765141606330872, - -0.9300433397293091, - -0.15196305513381958, - -0.2200203537940979, - -0.06051275506615639, - -0.6840062737464905, - -1.0964292287826538, - -0.17654964327812195, - -0.18547140061855316, - -0.06710249185562134, - -0.4758152365684509, - -0.6657928228378296, - -0.10342729091644287, - -0.10059614479541779, - -0.046978313475847244, - -0.410809725522995, - -0.428723007440567, - -0.06053968518972397, - -0.06518109142780304, - -0.030038274824619293, - -0.3271780014038086 + -3.010035514831543, + -0.011503085494041443, + -1.3661863803863525, + -0.7587734460830688, + -1.4640830755233765, + -1.4567692279815674, + -1.0451016426086426, + -2.3799993991851807, + -1.3697059154510498, + -1.0724036693572998, + -0.5331835150718689, + -1.8569976091384888, + -0.9731019735336304, + -0.05548504367470741, + -0.29222357273101807, + -0.2191229909658432, + -0.23294074833393097, + -0.6115235090255737, + -0.39632147550582886, + -1.1302311420440674, + -0.3844905197620392, + -0.7445144653320312, + -0.20952190458774567, + -0.29755353927612305, + -0.0599716454744339, + -0.008489826694130898, + -0.03623323515057564 ] }, - "mem-max-allocated-bytes": 53350692864, - "lifetime_prefill_token_count": 88 + "mem-max-allocated-bytes": 48828598272, + "lifetime_prefill_token_count": 88, + "async_sched_step_count": 0, + "async_sched_compaction_step_count": 0 } \ No newline at end of file diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_flashinfer/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_flashinfer/model_config.yaml index e989be22f7e..2491bd21433 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_flashinfer/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_flashinfer/model_config.yaml @@ -74,3 +74,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_mamba_bf16_states/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_mamba_bf16_states/model_config.yaml index 9affc9878f9..2a3de19fab6 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_mamba_bf16_states/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_dynamic_inference_tp1_pp1_dp8_583m_mamba_bf16_states/model_config.yaml @@ -73,3 +73,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_flextron_nightly_tp2_pp1_ep2_dgx_h100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_flextron_nightly_tp2_pp1_ep2_dgx_h100_1N8G/model_config.yaml index 493aaa31b16..34e4c07817b 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_flextron_nightly_tp2_pp1_ep2_dgx_h100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_flextron_nightly_tp2_pp1_ep2_dgx_h100_1N8G/model_config.yaml @@ -125,3 +125,4 @@ MODEL_ARGS: --slice: true --router-std: 0.1 TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp1_cp1_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp1_cp1_dgx_a100_1N8G/model_config.yaml index 9add53f8a49..b451a9ad5a5 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp1_cp1_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp1_cp1_dgx_a100_1N8G/model_config.yaml @@ -58,3 +58,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp2_vpp2_cp1_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp2_vpp2_cp1_dgx_a100_1N8G/model_config.yaml index 25df6aa0359..a5d76369b4b 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp2_vpp2_cp1_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp2_vpp2_cp1_dgx_a100_1N8G/model_config.yaml @@ -59,3 +59,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp4_cp1_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp4_cp1_dgx_a100_1N8G/model_config.yaml index fe4f9e63714..08707cb1d2c 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp4_cp1_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp1_pp4_cp1_dgx_a100_1N8G/model_config.yaml @@ -60,3 +60,4 @@ MODEL_ARGS: --async-strategy: mcore --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp1_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp1_dgx_a100_1N8G/model_config.yaml index 2339f7a7ce9..b4593e6102d 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp1_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp1_dgx_a100_1N8G/model_config.yaml @@ -56,3 +56,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp4_dgx_a100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp4_dgx_a100_1N8G/model_config.yaml index 3efc155949f..b203b225784 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp4_dgx_a100_1N8G/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_mr_mcore_te_tp2_pp1_cp4_dgx_a100_1N8G/model_config.yaml @@ -60,3 +60,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_nemotron_v3_pico_7b_a1b_tp1_ep8_QAD_dgx_h100_1N8G/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/hybrid/hybrid_nemotron_v3_pico_7b_a1b_tp1_ep8_QAD_dgx_h100_1N8G/golden_values_dev_dgx_h100.json new file mode 100644 index 00000000000..2922120bca8 --- /dev/null +++ b/tests/functional_tests/test_cases/hybrid/hybrid_nemotron_v3_pico_7b_a1b_tp1_ep8_QAD_dgx_h100_1N8G/golden_values_dev_dgx_h100.json @@ -0,0 +1,36 @@ +{ + "total loss": { + "start_step": 1, + "end_step": 10, + "step_interval": 1, + "values": { + "1": 1.58285, + "2": 1.6902, + "3": 0.08676, + "4": 0.064, + "5": 0.06231, + "6": 0.06777, + "7": 0.16431, + "8": 0.06122, + "9": 0.06703, + "10": 0.05211 + } + }, + "lm loss": { + "start_step": 1, + "end_step": 10, + "step_interval": 1, + "values": { + "1": 0.0, + "2": 0.0, + "3": 0.0, + "4": 0.0, + "5": 0.0, + "6": 0.0, + "7": 0.0, + "8": 0.0, + "9": 0.0, + "10": 0.0 + } + } +} diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_nemotron_v3_pico_7b_a1b_tp1_ep8_QAD_dgx_h100_1N8G/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_nemotron_v3_pico_7b_a1b_tp1_ep8_QAD_dgx_h100_1N8G/model_config.yaml new file mode 100644 index 00000000000..8e9946adcca --- /dev/null +++ b/tests/functional_tests/test_cases/hybrid/hybrid_nemotron_v3_pico_7b_a1b_tp1_ep8_QAD_dgx_h100_1N8G/model_config.yaml @@ -0,0 +1,191 @@ +ENV_VARS: + CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 + NCCL_ALGO: Ring + CUBLAS_WORKSPACE_CONFIG: ":4096:8" + TRITON_CACHE_AUTOTUNING: 0 + MAMBA_DETERMINISTIC: 1 + # Paths + MODEL_BF16_CKPT: "${DATA_PATH}/model/nemotron_v3_pico_7b-a1b/3T-token_deeparch" + PTQ_QUANTIZED_CKPT: "${DATA_CACHE_PATH}/ptq_quantized_ckpt" + TOKENIZER: "nvidia/NVIDIA-Nemotron-3-Nano-30B-A3B-BF16" + HF_HOME: "${DATA_PATH}/hf_home" +BEFORE_SCRIPT: | + # Stage 1: PTQ quantization via quantize.py directly + echo -e "\n=== Stage 1: Running PTQ (NVFP4) ===\n" + cd /opt/megatron-lm + # Env vars that arguments.sh would normally set + export TOKENIZERS_PARALLELISM=False + export OMP_NUM_THREADS=1 + export NCCL_IB_SL=1 + export NCCL_IB_TIMEOUT=22 + uv run --no-sync python -m torch.distributed.run --nproc_per_node=8 \ + examples/post_training/modelopt/quantize.py \ + --deterministic-mode \ + --micro-batch-size 1 \ + --save-interval 100000 \ + --bf16 \ + --seq-length 4096 \ + --max-position-embeddings 4096 \ + --tokenizer-type HuggingFaceTokenizer \ + --tokenizer-model ${TOKENIZER} \ + --tensor-model-parallel-size 1 \ + --expert-model-parallel-size 8 \ + --expert-tensor-parallel-size 1 \ + --pipeline-model-parallel-size 1 \ + --context-parallel-size 1 \ + --hidden-size 1216 \ + --num-attention-heads 32 \ + --group-query-attention \ + --num-query-groups 2 \ + --ffn-hidden-size 896 \ + --kv-channels 128 \ + --squared-relu \ + --normalization RMSNorm \ + --disable-bias-linear \ + --attention-dropout 0.0 \ + --hidden-dropout 0.0 \ + --position-embedding-type none \ + --untie-embeddings-and-output-weights \ + --init-method-std 0.0256 \ + --hybrid-layer-pattern MEMEM*EMEMEM*EMEMEM*EMEMEM*EMEMEM*EMEMEMEM*EMEMEMEME \ + --mamba-num-heads 64 \ + --export-model-type MambaModel \ + --num-experts 128 \ + --moe-router-topk 6 \ + --moe-aux-loss-coeff 1e-4 \ + --moe-router-topk-scaling-factor 2.5 \ + --moe-router-enable-expert-bias \ + --moe-router-dtype fp32 \ + --moe-router-score-function sigmoid \ + --moe-router-load-balancing-type seq_aux_loss \ + --moe-shared-expert-intermediate-size 3712 \ + --moe-token-dispatcher-type alltoall \ + --moe-grouped-gemm \ + --use-fused-weighted-squared-relu \ + --attention-backend fused \ + --disable-gloo-process-groups \ + --no-create-attention-mask-in-dataloader \ + --ckpt-format torch_dist \ + --ckpt-fully-parallel-load \ + --load ${MODEL_BF16_CKPT} \ + --save ${PTQ_QUANTIZED_CKPT} \ + --finetune \ + --auto-detect-ckpt-format \ + --distributed-timeout-minutes 30 \ + --export-quant-cfg MAMBA_MOE_NVFP4_CONSERVATIVE_CFG \ + --calib-dataset-path-or-name cnn_dailymail \ + --calib-size 256 \ + --calib-batch-size 8 \ + --skip-generate \ + --export-te-mcore-model + #--export-default-te-spec # TODO(aanoosheh): undo once TE fix is released + echo -e "\n=== Stage 1 complete ===\n" +MODEL_ARGS: + # KD teacher/config + --export-te-mcore-model: true + #--export-default-te-spec: true # TODO(aanoosheh): undo once TE fix is released + --export-kd-teacher-load: ${MODEL_BF16_CKPT} + --auto-detect-ckpt-format: true + --finetune: true + # Architecture + --hidden-size: 1216 + --num-attention-heads: 32 + --group-query-attention: true + --num-query-groups: 2 + --ffn-hidden-size: 896 + --kv-channels: 128 + --squared-relu: true + --normalization: RMSNorm + --disable-bias-linear: true + --attention-dropout: 0.0 + --hidden-dropout: 0.0 + --position-embedding-type: none + --untie-embeddings-and-output-weights: true + --init-method-std: 0.0256 + --hybrid-layer-pattern: MEMEM*EMEMEM*EMEMEM*EMEMEM*EMEMEM*EMEMEMEM*EMEMEMEME + --mamba-num-heads: 64 + --export-model-type: MambaModel + # MoE + --num-experts: 128 + --moe-router-topk: 6 + --moe-aux-loss-coeff: 1e-4 + --moe-router-topk-scaling-factor: 2.5 + --moe-router-enable-expert-bias: true + --moe-router-dtype: fp32 + --moe-router-score-function: sigmoid + --moe-router-load-balancing-type: seq_aux_loss + --moe-shared-expert-intermediate-size: 3712 + --moe-token-dispatcher-type: alltoall + --moe-grouped-gemm: true + --use-fused-weighted-squared-relu: true + # Tokenizer + --tokenizer-type: SFTTokenizer + --tokenizer-model: ${TOKENIZER} + --sft: true + --sft-tokenizer-prompt-format: identity + --bf16: true + # Parallelism + --tensor-model-parallel-size: 1 # TODO(aanoosheh): can change to 2 once TE fix is released + --expert-model-parallel-size: 4 + --expert-tensor-parallel-size: 1 + --pipeline-model-parallel-size: 1 + --context-parallel-size: 2 + --sequence-parallel: true + # Infrastructure + --attention-backend: fused + --disable-gloo-process-groups: true + --no-create-attention-mask-in-dataloader: true + --ddp-num-buckets: 8 + --override-opt_param-scheduler: true + --num-workers: 1 + --ckpt-format: torch_dist + --ckpt-fully-parallel-save: true + --ckpt-fully-parallel-load: true + --ckpt-assume-constant-structure: true + # Training + --micro-batch-size: 1 + --global-batch-size: 8 + --seq-length: 2048 # TODO(aanoosheh): change to 4096 once TE fix is released + --max-position-embeddings: 2048 # TODO(aanoosheh): change to 4096 once TE fix is released + --train-iters: 10 + --lr: 0.00015 + --lr-decay-style: cosine + --lr-decay-iters: 320000 + --min-lr: 1.0e-5 + --weight-decay: 1e-2 + --clip-grad: 1.0 + --lr-warmup-fraction: .01 + --use-distributed-optimizer: true + --overlap-param-gather: true + --overlap-grad-reduce: true + # Checkpoint & data + --save: ${CHECKPOINT_SAVE_PATH} + --load: ${PTQ_QUANTIZED_CKPT} + --data-path: "${DATA_PATH}/text/nemotron-3-super-sft_train-sample.jsonl" + --split: "949,50,1" + --distributed-backend: nccl + --transformer-impl: transformer_engine + --data-cache-path: ${DATA_CACHE_PATH} + # Logging + --log-interval: 1 + --save-interval: 10 + --eval-interval: 10 + --eval-iters: 2 + --log-params-norm: true + --log-num-zeros-in-grad: true + --log-validation-ppl-to-tensorboard: true + --log-timers-to-tensorboard: true + --tensorboard-dir: ${TENSORBOARD_PATH} + --log-memory-to-tensorboard: true + --timing-log-level: 0 + --no-gradient-accumulation-fusion: true + --distributed-timeout-minutes: 30 + # Etc + --deterministic-mode: true + --exit-interval: 10 +TEST_TYPE: regular +METRICS: + - lm loss + - total loss +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_cudagraphs/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_cudagraphs/model_config.yaml index 02c5cc3055c..9db4d9e247d 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_cudagraphs/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_cudagraphs/model_config.yaml @@ -72,3 +72,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_logitsmatch/model_config.yaml index 2543f59e668..75f341931a0 100644 --- a/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/hybrid/hybrid_static_inference_tp1_pp1_2B_logitsmatch/model_config.yaml @@ -68,3 +68,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp1_dp8/model_config.yaml b/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp1_dp8/model_config.yaml index e95856e7308..4aecf7bc60a 100644 --- a/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp1_dp8/model_config.yaml +++ b/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp1_dp8/model_config.yaml @@ -58,3 +58,4 @@ METRICS: - "num-zeros" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp1_dp8_seq_packing/model_config.yaml b/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp1_dp8_seq_packing/model_config.yaml index 2e86278fa67..89a71c25d9b 100644 --- a/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp1_dp8_seq_packing/model_config.yaml +++ b/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp1_dp8_seq_packing/model_config.yaml @@ -63,3 +63,4 @@ METRICS: - "num-zeros" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp2_dp8/model_config.yaml b/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp2_dp8/model_config.yaml index 37c55e4cd93..43acdf5cce9 100644 --- a/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp2_dp8/model_config.yaml +++ b/tests/functional_tests/test_cases/mimo/mimo_vlm_pretrain_convergence_tp1_pp1_cp2_dp8/model_config.yaml @@ -63,3 +63,4 @@ METRICS: - "num-zeros" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2/golden_values_dev_dgx_gb200.json b/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2/golden_values_dev_dgx_gb200.json index 857dfc6f69e..868798f915e 100644 --- a/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2/golden_values_dev_dgx_gb200.json +++ b/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2/golden_values_dev_dgx_gb200.json @@ -4,56 +4,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 10.90136, - "2": 10.93243, - "3": 10.64755, - "4": 10.41183, - "5": 10.40045, - "6": 10.36588, - "7": 10.11237, - "8": 9.90152, - "9": 9.96409, - "10": 9.51308, - "11": 10.16314, - "12": 9.86212, - "13": 9.8691, - "14": 9.90016, - "15": 9.50636, - "16": 9.4505, - "17": 9.25987, - "18": 9.30236, - "19": 9.20973, - "20": 8.97225, - "21": 9.00508, - "22": 8.60641, - "23": 9.09405, - "24": 8.68494, - "25": 8.50601, - "26": 8.73822, - "27": 8.82506, - "28": 8.95591, - "29": 8.94408, - "30": 8.49401, - "31": 7.96274, - "32": 8.67873, - "33": 8.74695, - "34": 8.26598, - "35": 8.35396, - "36": 8.27, - "37": 8.507, - "38": 8.2667, - "39": 8.63224, - "40": 8.22542, - "41": 8.27217, - "42": 8.42381, - "43": 8.00559, - "44": 8.11452, - "45": 7.99742, - "46": 8.08286, - "47": 8.38391, - "48": 8.09449, - "49": 7.72776, - "50": 8.19243 + "1": 10.90535, + "2": 10.912, + "3": 10.35928, + "4": 10.0822, + "5": 9.85592, + "6": 9.51097, + "7": 9.37079, + "8": 9.20032, + "9": 9.09294, + "10": 8.9534, + "11": 8.99361, + "12": 8.85188, + "13": 8.84429, + "14": 8.77849, + "15": 8.54032, + "16": 8.69694, + "17": 8.37147, + "18": 8.39638, + "19": 8.27818, + "20": 8.27646, + "21": 8.14342, + "22": 8.13757, + "23": 8.16585, + "24": 8.15514, + "25": 8.19365, + "26": 7.90145, + "27": 7.98979, + "28": 8.05491, + "29": 8.10553, + "30": 7.83891, + "31": 8.0464, + "32": 7.95497, + "33": 7.86316, + "34": 7.87278, + "35": 7.60145, + "36": 7.9792, + "37": 7.77922, + "38": 7.5885, + "39": 7.76987, + "40": 7.57503, + "41": 7.75856, + "42": 7.70609, + "43": 7.76788, + "44": 7.60289, + "45": 7.59598, + "46": 7.6509, + "47": 7.71542, + "48": 7.65185, + "49": 7.61001, + "50": 7.58617 } }, "num-zeros": { @@ -61,56 +61,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 30093584.0, - "2": 30296868.0, - "3": 29948216.0, - "4": 30610452.0, - "5": 30079954.0, - "6": 30406440.0, - "7": 30137332.0, - "8": 30305308.0, - "9": 30210788.0, - "10": 30295208.0, - "11": 29844004.0, - "12": 29800708.0, - "13": 30292576.0, - "14": 29729006.0, - "15": 30191246.0, - "16": 30201442.0, - "17": 30181084.0, - "18": 29934542.0, - "19": 29972104.0, - "20": 30057016.0, - "21": 30106920.0, - "22": 30162128.0, - "23": 29890266.0, - "24": 30139388.0, - "25": 30186512.0, - "26": 29894670.0, - "27": 29809340.0, - "28": 29798496.0, - "29": 29879152.0, - "30": 29980394.0, - "31": 30332976.0, - "32": 29938440.0, - "33": 29910984.0, - "34": 30200552.0, - "35": 30158306.0, - "36": 29940436.0, - "37": 29844712.0, - "38": 30271032.0, - "39": 30169892.0, - "40": 30020722.0, - "41": 30016120.0, - "42": 30027286.0, - "43": 30358212.0, - "44": 30109930.0, - "45": 30023796.0, - "46": 30256184.0, - "47": 29987288.0, - "48": 30298758.0, - "49": 30082702.0, - "50": 30275984.0 + "1": 43613800.0, + "2": 43061692.0, + "3": 43018972.0, + "4": 42752120.0, + "5": 43058592.0, + "6": 43163140.0, + "7": 43319396.0, + "8": 43311548.0, + "9": 43217716.0, + "10": 43253296.0, + "11": 43506680.0, + "12": 43393248.0, + "13": 43101692.0, + "14": 43082760.0, + "15": 43206672.0, + "16": 43007424.0, + "17": 43331104.0, + "18": 43052920.0, + "19": 43409776.0, + "20": 43058788.0, + "21": 43163380.0, + "22": 42934560.0, + "23": 43641040.0, + "24": 43096052.0, + "25": 43026180.0, + "26": 43692032.0, + "27": 43430164.0, + "28": 43021920.0, + "29": 43109872.0, + "30": 43258052.0, + "31": 43137200.0, + "32": 43371032.0, + "33": 43064356.0, + "34": 42892228.0, + "35": 43263720.0, + "36": 43096816.0, + "37": 43344644.0, + "38": 43566448.0, + "39": 43389520.0, + "40": 43401520.0, + "41": 43231292.0, + "42": 43148956.0, + "43": 43157596.0, + "44": 43367328.0, + "45": 43323048.0, + "46": 43013076.0, + "47": 43078092.0, + "48": 43305572.0, + "49": 43418436.0, + "50": 43533312.0 } }, "mem-allocated-bytes": { @@ -118,56 +118,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 628048896.0, - "2": 628050432.0, - "3": 628050432.0, - "4": 628050432.0, - "5": 628050432.0, - "6": 628050432.0, - "7": 628050432.0, - "8": 628050432.0, - "9": 628050432.0, - "10": 628050432.0, - "11": 628050432.0, - "12": 628050432.0, - "13": 628050432.0, - "14": 628050432.0, - "15": 628050432.0, - "16": 628050432.0, - "17": 628050432.0, - "18": 628050432.0, - "19": 628050432.0, - "20": 628050432.0, - "21": 628050432.0, - "22": 628050432.0, - "23": 628050432.0, - "24": 628050432.0, - "25": 628050432.0, - "26": 628050432.0, - "27": 628050432.0, - "28": 628050432.0, - "29": 628050432.0, - "30": 628050432.0, - "31": 628050432.0, - "32": 628050432.0, - "33": 628050432.0, - "34": 628050432.0, - "35": 628050432.0, - "36": 628050432.0, - "37": 628050432.0, - "38": 628050432.0, - "39": 628050432.0, - "40": 628050432.0, - "41": 628050432.0, - "42": 628050432.0, - "43": 628050432.0, - "44": 628050432.0, - "45": 628050432.0, - "46": 628050432.0, - "47": 628050432.0, - "48": 628050432.0, - "49": 628050432.0, - "50": 628050432.0 + "1": 824764928.0, + "2": 824766464.0, + "3": 824766464.0, + "4": 824766464.0, + "5": 824766464.0, + "6": 824766464.0, + "7": 824766464.0, + "8": 824766464.0, + "9": 824766464.0, + "10": 824766464.0, + "11": 824766464.0, + "12": 824766464.0, + "13": 824766464.0, + "14": 824766464.0, + "15": 824766464.0, + "16": 824766464.0, + "17": 824766464.0, + "18": 824766464.0, + "19": 824766464.0, + "20": 824766464.0, + "21": 824766464.0, + "22": 824766464.0, + "23": 824766464.0, + "24": 824766464.0, + "25": 824766464.0, + "26": 824766464.0, + "27": 824766464.0, + "28": 824766464.0, + "29": 824766464.0, + "30": 824766464.0, + "31": 824766464.0, + "32": 824766464.0, + "33": 824766464.0, + "34": 824766464.0, + "35": 824766464.0, + "36": 824766464.0, + "37": 824766464.0, + "38": 824766464.0, + "39": 824766464.0, + "40": 824766464.0, + "41": 824766464.0, + "42": 824766464.0, + "43": 824766464.0, + "44": 824766464.0, + "45": 824766464.0, + "46": 824766464.0, + "47": 824766464.0, + "48": 824766464.0, + "49": 824766464.0, + "50": 824766464.0 } }, "mem-max-allocated-bytes": { @@ -175,56 +175,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 2302537728.0, - "2": 2346255360.0, - "3": 2354982400.0, - "4": 2357787648.0, - "5": 2360281088.0, - "6": 2360281088.0, - "7": 2360281088.0, - "8": 2360281088.0, - "9": 2360281088.0, - "10": 2360281088.0, - "11": 2360281088.0, - "12": 2360281088.0, - "13": 2360281088.0, - "14": 2361527808.0, - "15": 2361527808.0, - "16": 2361527808.0, - "17": 2361527808.0, - "18": 2361527808.0, - "19": 2361527808.0, - "20": 2361527808.0, - "21": 2361527808.0, - "22": 2361527808.0, - "23": 2361527808.0, - "24": 2361527808.0, - "25": 2361527808.0, - "26": 2361527808.0, - "27": 2361527808.0, - "28": 2361527808.0, - "29": 2361527808.0, - "30": 2361527808.0, - "31": 2361527808.0, - "32": 2361527808.0, - "33": 2361527808.0, - "34": 2361527808.0, - "35": 2361527808.0, - "36": 2361527808.0, - "37": 2361527808.0, - "38": 2361527808.0, - "39": 2369631744.0, - "40": 2369631744.0, - "41": 2369631744.0, - "42": 2369631744.0, - "43": 2369631744.0, - "44": 2369631744.0, - "45": 2376800768.0, - "46": 2376800768.0, - "47": 2376800768.0, - "48": 2376800768.0, - "49": 2376800768.0, - "50": 2376800768.0 + "1": 33842712576.0, + "2": 33868759040.0, + "3": 33868759040.0, + "4": 33868759040.0, + "5": 34014625792.0, + "6": 34014625792.0, + "7": 34014625792.0, + "8": 34014625792.0, + "9": 34014625792.0, + "10": 34014625792.0, + "11": 34308540416.0, + "12": 34308540416.0, + "13": 34308540416.0, + "14": 34308540416.0, + "15": 34308540416.0, + "16": 34308540416.0, + "17": 34308540416.0, + "18": 34308540416.0, + "19": 34308540416.0, + "20": 34308540416.0, + "21": 34308540416.0, + "22": 34308540416.0, + "23": 34308540416.0, + "24": 34308540416.0, + "25": 34308540416.0, + "26": 34308540416.0, + "27": 34308540416.0, + "28": 34308540416.0, + "29": 34308540416.0, + "30": 34308540416.0, + "31": 34308540416.0, + "32": 34308540416.0, + "33": 34308540416.0, + "34": 34308540416.0, + "35": 34308540416.0, + "36": 34308540416.0, + "37": 34308540416.0, + "38": 34308540416.0, + "39": 34308540416.0, + "40": 34308540416.0, + "41": 34308540416.0, + "42": 34308540416.0, + "43": 34308540416.0, + "44": 34308540416.0, + "45": 34308540416.0, + "46": 34308540416.0, + "47": 34308540416.0, + "48": 34308540416.0, + "49": 34308540416.0, + "50": 34308540416.0 } }, "iteration-time": { @@ -232,56 +232,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": "nan", - "2": 12.8321, - "3": 0.12373, - "4": 0.08346, - "5": 0.08416, - "6": 0.07239, - "7": 0.08224, - "8": 0.07207, - "9": 0.08251, - "10": 0.08562, - "11": 0.08731, - "12": 0.08106, - "13": 0.07755, - "14": 0.07683, - "15": 0.08155, - "16": 0.07425, - "17": 0.07476, - "18": 0.07543, - "19": 0.07529, - "20": 0.07404, - "21": 0.07799, - "22": 0.07111, - "23": 0.07327, - "24": 0.07359, - "25": 0.07084, - "26": 0.07257, - "27": 0.07256, - "28": 0.07299, - "29": 0.07575, - "30": 0.07277, - "31": 0.0723, - "32": 0.0756, - "33": 0.07168, - "34": 0.0718, - "35": 0.07182, - "36": 0.07978, - "37": 0.07074, - "38": 0.07522, - "39": 0.07461, - "40": 0.07262, - "41": 0.07271, - "42": 0.07214, - "43": 0.07172, - "44": 0.07212, - "45": 0.07086, - "46": 0.07628, - "47": 0.07302, - "48": 0.07252, - "49": 0.07203, - "50": 0.07695 + "1": 0.0, + "2": 15.6569, + "3": 0.33059, + "4": 0.27341, + "5": 0.28876, + "6": 0.26934, + "7": 0.2514, + "8": 0.25469, + "9": 0.25137, + "10": 0.24816, + "11": 0.24194, + "12": 0.24407, + "13": 0.24915, + "14": 0.24415, + "15": 0.25581, + "16": 0.24996, + "17": 0.25951, + "18": 0.25195, + "19": 0.24463, + "20": 0.24788, + "21": 0.25507, + "22": 0.25148, + "23": 0.24187, + "24": 0.24372, + "25": 0.23858, + "26": 0.24133, + "27": 0.25301, + "28": 0.25271, + "29": 0.25946, + "30": 0.25355, + "31": 0.25977, + "32": 0.24357, + "33": 0.25322, + "34": 0.25063, + "35": 0.23954, + "36": 0.23716, + "37": 0.24081, + "38": 0.24199, + "39": 0.23279, + "40": 0.23709, + "41": 0.23498, + "42": 0.23858, + "43": 0.23672, + "44": 0.24102, + "45": 0.23953, + "46": 0.23516, + "47": 0.23748, + "48": 0.24031, + "49": 0.23794, + "50": 0.23369 } } } \ No newline at end of file diff --git a/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2/model_config.yaml b/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2/model_config.yaml index 7b820614bf8..8a5fe85b09a 100644 --- a/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2/model_config.yaml @@ -1,4 +1,4 @@ -# DeepSeek-style proxy: 2 layers — first dense, second MoE with expert parallelism 4. +# DeepSeek-style proxy: 8 layers — 3 dense, Rest MoE with expert parallelism 2. # Used for functional testing of dense+MoE hybrid and EP. ENV_VARS: NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 @@ -42,8 +42,8 @@ MODEL_ARGS: --use-mcore-models: true --sequence-parallel: true --disable-bias-linear: true - --micro-batch-size: 4 - --global-batch-size: 32 + --micro-batch-size: 32 + --global-batch-size: 256 --train-iters: 50 --exit-duration-in-mins: 60 --no-check-for-nan-in-loss-and-grad: true @@ -62,7 +62,7 @@ MODEL_ARGS: # Network: 2 layers — first dense, second MoE - --num-layers: 2 + --num-layers: 8 --hidden-size: 512 --ffn-hidden-size: 2048 --num-attention-heads: 8 @@ -92,7 +92,7 @@ MODEL_ARGS: --adam-beta2: 0.95 # MoE args (DeepSeek-style): 1 dense layer, 1 MoE layer, ep=4 --num-experts: 8 - --moe-layer-freq: ([0]*1+[1]*1) + --moe-layer-freq: ([0]*3+[1]*5) --moe-ffn-hidden-size: 1024 --moe-shared-expert-intermediate-size: 1024 --moe-router-load-balancing-type: seq_aux_loss @@ -143,3 +143,4 @@ METRICS: - "lm loss" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2_1node/model_config.yaml b/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2_1node/model_config.yaml index a243bcf6b84..9cf511795a9 100644 --- a/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2_1node/model_config.yaml @@ -139,3 +139,4 @@ METRICS: - "lm loss" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2_ep_overlap/model_config.yaml b/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2_ep_overlap/model_config.yaml index 0d1e04af73d..bac8829e669 100644 --- a/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2_ep_overlap/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/deepseek_proxy_fsdp_ep2_fsdp2_ep_overlap/model_config.yaml @@ -133,3 +133,4 @@ METRICS: - "lm loss" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_cp2_pp2_ep2_te_4experts2parallel_nondeterministic/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_cp2_pp2_ep2_te_4experts2parallel_nondeterministic/model_config.yaml index b3e8a82de72..b8ec491215c 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_cp2_pp2_ep2_te_4experts2parallel_nondeterministic/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_cp2_pp2_ep2_te_4experts2parallel_nondeterministic/model_config.yaml @@ -59,3 +59,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_cp2_pp2_ep2_te_4experts2parallel_nondeterministic_dp_last/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_cp2_pp2_ep2_te_4experts2parallel_nondeterministic_dp_last/model_config.yaml index 59887d4eec9..70f522d1d4c 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_cp2_pp2_ep2_te_4experts2parallel_nondeterministic_dp_last/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_cp2_pp2_ep2_te_4experts2parallel_nondeterministic_dp_last/model_config.yaml @@ -61,3 +61,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp1_pp1_te_4experts_groupedGEMM_op_fuser/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp1_pp1_te_4experts_groupedGEMM_op_fuser/model_config.yaml index 4a4fc48fb8b..14352931d87 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp1_pp1_te_4experts_groupedGEMM_op_fuser/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp1_pp1_te_4experts_groupedGEMM_op_fuser/model_config.yaml @@ -61,3 +61,4 @@ MODEL_ARGS: --no-bias-gelu-fusion: true --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_te_8experts2parallel_dist_optimizer/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_te_8experts2parallel_dist_optimizer/model_config.yaml index c4bc4528090..4038a4e55e8 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_te_8experts2parallel_dist_optimizer/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_te_8experts2parallel_dist_optimizer/model_config.yaml @@ -63,3 +63,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: frozen-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_te_8experts2parallel_groupedGEMM/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_te_8experts2parallel_groupedGEMM/model_config.yaml index bdef7c88323..763a92f7b64 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_te_8experts2parallel_groupedGEMM/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_frozen_resume_torch_dist_te_8experts2parallel_groupedGEMM/model_config.yaml @@ -59,3 +59,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: frozen-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_dist_optimizer/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_dist_optimizer/model_config.yaml index 2a8a2a5d72b..51b9d2e103a 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_dist_optimizer/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_dist_optimizer/model_config.yaml @@ -64,3 +64,4 @@ MODEL_ARGS: --use-persistent-ckpt-worker: true --verify-integrity: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_dist_optimizer_1node/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_dist_optimizer_1node/model_config.yaml index 764c576645e..f8f7222f004 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_dist_optimizer_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_dist_optimizer_1node/model_config.yaml @@ -63,3 +63,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_groupedGEMM/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_groupedGEMM/model_config.yaml index 381039b4905..d141f68132d 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_groupedGEMM/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_groupedGEMM/model_config.yaml @@ -59,3 +59,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_top2router/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_top2router/model_config.yaml index 86a775807b8..2c11ccb1acb 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_top2router/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_resume_torch_dist_te_8experts2parallel_top2router/model_config.yaml @@ -61,3 +61,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective/golden_values_dev_dgx_h100.json index 8a7452aa143..3eefe1b2bb3 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective/golden_values_dev_dgx_h100.json +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective/golden_values_dev_dgx_h100.json @@ -6,54 +6,54 @@ "values": { "1": 10.92563, "2": 10.91638, - "3": 10.92433, - "4": 10.93217, - "5": 10.92999, - "6": 10.92662, - "7": 10.92571, - "8": 10.92333, - "9": 10.92825, - "10": 10.91605, - "11": 10.91854, - "12": 10.92399, - "13": 10.91037, - "14": 10.90685, - "15": 10.90136, - "16": 10.88661, - "17": 10.88849, - "18": 10.88662, - "19": 10.8857, - "20": 10.83765, - "21": 10.82761, - "22": 10.81538, - "23": 10.8078, - "24": 10.78018, - "25": 10.778, - "26": 10.76109, - "27": 10.74912, - "28": 10.69195, - "29": 10.66617, - "30": 10.63122, - "31": 10.62222, - "32": 10.61543, - "33": 10.57901, - "34": 10.54677, - "35": 10.54608, - "36": 10.53419, - "37": 10.50598, - "38": 10.50458, - "39": 10.47274, - "40": 10.45062, - "41": 10.42727, - "42": 10.41444, - "43": 10.40126, - "44": 10.3705, - "45": 10.38167, - "46": 10.33539, - "47": 10.32458, - "48": 10.28718, - "49": 10.28599, - "50": 10.27739 + "3": 10.92459, + "4": 10.93245, + "5": 10.93019, + "6": 10.92651, + "7": 10.92566, + "8": 10.92291, + "9": 10.92842, + "10": 10.91709, + "11": 10.9186, + "12": 10.92364, + "13": 10.91032, + "14": 10.906, + "15": 10.90106, + "16": 10.88602, + "17": 10.88857, + "18": 10.8863, + "19": 10.88563, + "20": 10.83831, + "21": 10.82801, + "22": 10.81606, + "23": 10.80814, + "24": 10.77971, + "25": 10.77746, + "26": 10.76123, + "27": 10.74966, + "28": 10.69186, + "29": 10.66597, + "30": 10.63047, + "31": 10.62182, + "32": 10.61491, + "33": 10.5784, + "34": 10.54596, + "35": 10.54573, + "36": 10.53409, + "37": 10.5049, + "38": 10.50267, + "39": 10.47204, + "40": 10.44914, + "41": 10.426, + "42": 10.41331, + "43": 10.39913, + "44": 10.36866, + "45": 10.37972, + "46": 10.3331, + "47": 10.32171, + "48": 10.28474, + "49": 10.28349, + "50": 10.27386 } }, "num-zeros": { @@ -61,56 +61,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 18949.0, - "2": 19099.0, - "3": 19050.0, - "4": 19118.0, - "5": 19083.0, - "6": 18713.0, - "7": 18983.0, - "8": 18573.0, - "9": 19074.0, - "10": 19633.0, - "11": 18864.0, - "12": 18680.0, - "13": 18983.0, - "14": 19635.0, - "15": 18534.0, - "16": 18776.0, - "17": 19027.0, - "18": 19179.0, - "19": 19256.0, - "20": 19131.0, - "21": 18678.0, - "22": 19277.0, - "23": 19003.0, - "24": 19177.0, - "25": 18328.0, - "26": 18877.0, - "27": 19140.0, - "28": 18575.0, - "29": 18512.0, - "30": 18566.0, - "31": 18581.0, - "32": 19224.0, - "33": 18078.0, - "34": 18612.0, - "35": 18599.0, - "36": 18325.0, - "37": 18458.0, - "38": 18488.0, - "39": 18900.0, - "40": 19303.0, - "41": 18755.0, - "42": 18613.0, - "43": 19044.0, - "44": 18919.0, - "45": 20364.0, - "46": 19848.0, - "47": 20037.0, - "48": 19837.0, - "49": 21694.0, - "50": 20114.0 + "1": 36671.0, + "2": 37173.0, + "3": 36899.0, + "4": 37068.0, + "5": 36679.0, + "6": 35978.0, + "7": 37031.0, + "8": 36229.0, + "9": 37090.0, + "10": 37967.0, + "11": 36941.0, + "12": 36081.0, + "13": 36942.0, + "14": 37392.0, + "15": 36186.0, + "16": 36099.0, + "17": 36904.0, + "18": 36981.0, + "19": 37455.0, + "20": 36609.0, + "21": 36447.0, + "22": 36641.0, + "23": 36893.0, + "24": 37094.0, + "25": 35824.0, + "26": 36997.0, + "27": 36498.0, + "28": 35858.0, + "29": 36103.0, + "30": 35720.0, + "31": 36188.0, + "32": 36884.0, + "33": 36092.0, + "34": 36008.0, + "35": 36767.0, + "36": 35759.0, + "37": 36357.0, + "38": 35639.0, + "39": 36859.0, + "40": 37400.0, + "41": 36473.0, + "42": 36460.0, + "43": 37019.0, + "44": 36370.0, + "45": 39844.0, + "46": 38478.0, + "47": 38793.0, + "48": 38437.0, + "49": 42666.0, + "50": 39545.0 } }, "mem-allocated-bytes": { @@ -118,56 +118,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 1027089408.0, - "2": 1027091456.0, - "3": 1027087360.0, - "4": 1027088384.0, - "5": 1027091456.0, - "6": 1027091456.0, - "7": 1027088896.0, - "8": 1027092480.0, + "1": 1027088896.0, + "2": 1027090944.0, + "3": 1027086848.0, + "4": 1027087360.0, + "5": 1027091968.0, + "6": 1027089920.0, + "7": 1027088384.0, + "8": 1027091968.0, "9": 1027091968.0, - "10": 1027089408.0, + "10": 1027089920.0, "11": 1027089920.0, - "12": 1027092480.0, + "12": 1027091456.0, "13": 1027090944.0, - "14": 1027092480.0, - "15": 1027090432.0, - "16": 1027088384.0, - "17": 1027089408.0, - "18": 1027090944.0, - "19": 1027088384.0, - "20": 1027090432.0, - "21": 1027092480.0, + "14": 1027091456.0, + "15": 1027089920.0, + "16": 1027088896.0, + "17": 1027087872.0, + "18": 1027090432.0, + "19": 1027088896.0, + "20": 1027089408.0, + "21": 1027091456.0, "22": 1027089920.0, - "23": 1027093504.0, + "23": 1027094016.0, "24": 1027092480.0, "25": 1027089408.0, - "26": 1027090944.0, - "27": 1027087360.0, - "28": 1027090432.0, - "29": 1027090432.0, - "30": 1027089920.0, - "31": 1027089408.0, - "32": 1027093504.0, - "33": 1027094016.0, - "34": 1027093504.0, - "35": 1027085824.0, - "36": 1027087872.0, + "26": 1027089408.0, + "27": 1027086848.0, + "28": 1027090944.0, + "29": 1027089920.0, + "30": 1027090432.0, + "31": 1027090432.0, + "32": 1027092992.0, + "33": 1027092480.0, + "34": 1027091968.0, + "35": 1027085312.0, + "36": 1027086336.0, "37": 1027088896.0, "38": 1027089920.0, - "39": 1027088384.0, - "40": 1027091968.0, + "39": 1027087360.0, + "40": 1027090944.0, "41": 1027088384.0, - "42": 1027089408.0, + "42": 1027087872.0, "43": 1027087872.0, - "44": 1027091456.0, + "44": 1027088384.0, "45": 1027090432.0, "46": 1027086336.0, - "47": 1027088384.0, - "48": 1027087360.0, - "49": 1027087360.0, - "50": 1027089920.0 + "47": 1027086848.0, + "48": 1027086848.0, + "49": 1027086336.0, + "50": 1027089408.0 } }, "mem-max-allocated-bytes": { @@ -175,113 +175,112 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 3058326528.0, - "2": 3298517504.0, - "3": 3298517504.0, - "4": 3298517504.0, - "5": 3300747776.0, - "6": 3300747776.0, - "7": 3300747776.0, - "8": 3300747776.0, - "9": 3300747776.0, - "10": 3300747776.0, - "11": 3300747776.0, - "12": 3300747776.0, - "13": 3300747776.0, - "14": 3300747776.0, - "15": 3300747776.0, - "16": 3300747776.0, - "17": 3300747776.0, - "18": 3300747776.0, - "19": 3300747776.0, - "20": 3300747776.0, - "21": 3300747776.0, - "22": 3300747776.0, - "23": 3300747776.0, - "24": 3300747776.0, - "25": 3300747776.0, - "26": 3300747776.0, - "27": 3300747776.0, - "28": 3300747776.0, - "29": 3300747776.0, - "30": 3300747776.0, - "31": 3300747776.0, - "32": 3300747776.0, - "33": 3300747776.0, - "34": 3300872192.0, - "35": 3300872192.0, - "36": 3300872192.0, - "37": 3300872192.0, - "38": 3300872192.0, - "39": 3300872192.0, - "40": 3300872192.0, - "41": 3300872192.0, - "42": 3300872192.0, - "43": 3300872192.0, - "44": 3300872192.0, - "45": 3300872192.0, - "46": 3300872192.0, - "47": 3300872192.0, - "48": 3300872192.0, - "49": 3300872192.0, - "50": 3300872192.0 + "1": 3059096576.0, + "2": 3298776064.0, + "3": 3298776064.0, + "4": 3298776064.0, + "5": 3298776064.0, + "6": 3298776064.0, + "7": 3298776064.0, + "8": 3299136512.0, + "9": 3299136512.0, + "10": 3299136512.0, + "11": 3299136512.0, + "12": 3299136512.0, + "13": 3299320320.0, + "14": 3299397120.0, + "15": 3299397120.0, + "16": 3299397120.0, + "17": 3299397120.0, + "18": 3299397120.0, + "19": 3299397120.0, + "20": 3299397120.0, + "21": 3299397120.0, + "22": 3299397120.0, + "23": 3300247552.0, + "24": 3300247552.0, + "25": 3300247552.0, + "26": 3300247552.0, + "27": 3300247552.0, + "28": 3300247552.0, + "29": 3300247552.0, + "30": 3300247552.0, + "31": 3300247552.0, + "32": 3300247552.0, + "33": 3300247552.0, + "34": 3300554752.0, + "35": 3300554752.0, + "36": 3300554752.0, + "37": 3300554752.0, + "38": 3300554752.0, + "39": 3300554752.0, + "40": 3300554752.0, + "41": 3300554752.0, + "42": 3300554752.0, + "43": 3300554752.0, + "44": 3300554752.0, + "45": 3300554752.0, + "46": 3300554752.0, + "47": 3300554752.0, + "48": 3300554752.0, + "49": 3300554752.0, + "50": 3300554752.0 } }, "iteration-time": { - "start_step": 1, + "start_step": 2, "end_step": 50, "step_interval": 1, "values": { - "1": "nan", - "2": 7.37375, - "3": 0.25401, - "4": 0.23, - "5": 0.23156, - "6": 0.22618, - "7": 0.22033, - "8": 0.2124, - "9": 0.21458, - "10": 0.2112, - "11": 0.22058, - "12": 0.21214, - "13": 0.20964, - "14": 0.21773, - "15": 0.21046, - "16": 0.21558, - "17": 0.21724, - "18": 0.21042, - "19": 0.2121, - "20": 0.21156, - "21": 0.2121, - "22": 0.20983, - "23": 0.22142, - "24": 0.21088, - "25": 0.21096, - "26": 0.2105, - "27": 0.21223, - "28": 0.21432, - "29": 0.20728, - "30": 0.20861, - "31": 0.20793, - "32": 0.20812, - "33": 0.20817, - "34": 0.20922, - "35": 0.20912, - "36": 0.21051, - "37": 0.21278, - "38": 0.21391, - "39": 0.2131, - "40": 0.21335, - "41": 0.21205, - "42": 0.20975, - "43": 0.2117, - "44": 0.21456, - "45": 0.21588, - "46": 0.21062, - "47": 0.21618, - "48": 0.21235, - "49": 0.21609, - "50": 0.21536 + "2": 5.12545, + "3": 0.33231, + "4": 0.34056, + "5": 0.32602, + "6": 0.3278, + "7": 0.32667, + "8": 0.30797, + "9": 0.30763, + "10": 0.30676, + "11": 0.31256, + "12": 0.3065, + "13": 0.3017, + "14": 0.29823, + "15": 0.30407, + "16": 0.3043, + "17": 0.29941, + "18": 1.04497, + "19": 0.30398, + "20": 0.30088, + "21": 0.31032, + "22": 0.30474, + "23": 0.30372, + "24": 0.30325, + "25": 0.31021, + "26": 0.2994, + "27": 0.30871, + "28": 0.29653, + "29": 0.29361, + "30": 0.29596, + "31": 0.2957, + "32": 0.29902, + "33": 0.2982, + "34": 0.29551, + "35": 0.29378, + "36": 0.31272, + "37": 0.30903, + "38": 0.30702, + "39": 0.30131, + "40": 0.3076, + "41": 0.30128, + "42": 0.29866, + "43": 0.30149, + "44": 0.30224, + "45": 0.29648, + "46": 0.30138, + "47": 0.29993, + "48": 0.29673, + "49": 0.29749, + "50": 0.30226 } } -} \ No newline at end of file +} diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective/model_config.yaml index 8ced6e37a52..9353bea2933 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective/model_config.yaml @@ -64,3 +64,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective_1node/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective_1node/model_config.yaml index 8ced6e37a52..9353bea2933 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_ddp_average_in_collective_1node/model_config.yaml @@ -64,3 +64,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_dist_optimizer/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_dist_optimizer/model_config.yaml index ec59433dc1b..3a9d1ec9a78 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_dist_optimizer/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_dist_optimizer/model_config.yaml @@ -61,3 +61,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM/golden_values_dev_dgx_h100.json index 83ebb282949..8267b6a3f65 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM/golden_values_dev_dgx_h100.json +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM/golden_values_dev_dgx_h100.json @@ -6,54 +6,54 @@ "values": { "1": 10.92563, "2": 10.91638, - "3": 10.92433, - "4": 10.93217, - "5": 10.92999, - "6": 10.92662, - "7": 10.92571, - "8": 10.92333, - "9": 10.92825, - "10": 10.91605, - "11": 10.91854, - "12": 10.92399, - "13": 10.91037, - "14": 10.90685, - "15": 10.90136, - "16": 10.88661, - "17": 10.88849, - "18": 10.88662, - "19": 10.8857, - "20": 10.83765, - "21": 10.82761, - "22": 10.81538, - "23": 10.8078, - "24": 10.78018, - "25": 10.778, - "26": 10.76109, - "27": 10.74912, - "28": 10.69195, - "29": 10.66617, - "30": 10.63122, - "31": 10.62222, - "32": 10.61543, - "33": 10.57901, - "34": 10.54677, - "35": 10.54608, - "36": 10.53419, - "37": 10.50598, - "38": 10.50458, - "39": 10.47274, - "40": 10.45062, - "41": 10.42727, - "42": 10.41444, - "43": 10.40126, - "44": 10.3705, - "45": 10.38167, - "46": 10.33539, - "47": 10.32458, - "48": 10.28718, - "49": 10.28599, - "50": 10.27739 + "3": 10.92459, + "4": 10.93245, + "5": 10.93019, + "6": 10.92651, + "7": 10.92566, + "8": 10.92291, + "9": 10.92842, + "10": 10.91709, + "11": 10.9186, + "12": 10.92364, + "13": 10.91032, + "14": 10.906, + "15": 10.90106, + "16": 10.88602, + "17": 10.88857, + "18": 10.8863, + "19": 10.88563, + "20": 10.83831, + "21": 10.82801, + "22": 10.81606, + "23": 10.80814, + "24": 10.77971, + "25": 10.77746, + "26": 10.76123, + "27": 10.74966, + "28": 10.69186, + "29": 10.66597, + "30": 10.63047, + "31": 10.62182, + "32": 10.61491, + "33": 10.5784, + "34": 10.54596, + "35": 10.54573, + "36": 10.53409, + "37": 10.5049, + "38": 10.50267, + "39": 10.47204, + "40": 10.44914, + "41": 10.426, + "42": 10.41331, + "43": 10.39913, + "44": 10.36866, + "45": 10.37972, + "46": 10.3331, + "47": 10.32171, + "48": 10.28474, + "49": 10.28349, + "50": 10.27386 } }, "num-zeros": { @@ -61,56 +61,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 18949.0, - "2": 19099.0, - "3": 19050.0, - "4": 19118.0, - "5": 19083.0, - "6": 18713.0, - "7": 18983.0, - "8": 18573.0, - "9": 19074.0, - "10": 19633.0, - "11": 18864.0, - "12": 18680.0, - "13": 18983.0, - "14": 19635.0, - "15": 18534.0, - "16": 18776.0, - "17": 19027.0, - "18": 19179.0, - "19": 19256.0, - "20": 19131.0, - "21": 18678.0, - "22": 19277.0, - "23": 19003.0, - "24": 19177.0, - "25": 18328.0, - "26": 18877.0, - "27": 19140.0, - "28": 18575.0, - "29": 18512.0, - "30": 18566.0, - "31": 18581.0, - "32": 19224.0, - "33": 18078.0, - "34": 18612.0, - "35": 18599.0, - "36": 18325.0, - "37": 18458.0, - "38": 18488.0, - "39": 18900.0, - "40": 19303.0, - "41": 18755.0, - "42": 18613.0, - "43": 19044.0, - "44": 18919.0, - "45": 20364.0, - "46": 19848.0, - "47": 20037.0, - "48": 19837.0, - "49": 21694.0, - "50": 20114.0 + "1": 36671.0, + "2": 37173.0, + "3": 36899.0, + "4": 37068.0, + "5": 36679.0, + "6": 35978.0, + "7": 37031.0, + "8": 36229.0, + "9": 37090.0, + "10": 37967.0, + "11": 36941.0, + "12": 36081.0, + "13": 36942.0, + "14": 37392.0, + "15": 36186.0, + "16": 36099.0, + "17": 36904.0, + "18": 36981.0, + "19": 37455.0, + "20": 36609.0, + "21": 36447.0, + "22": 36641.0, + "23": 36893.0, + "24": 37094.0, + "25": 35824.0, + "26": 36997.0, + "27": 36498.0, + "28": 35858.0, + "29": 36103.0, + "30": 35720.0, + "31": 36188.0, + "32": 36884.0, + "33": 36092.0, + "34": 36008.0, + "35": 36767.0, + "36": 35759.0, + "37": 36357.0, + "38": 35639.0, + "39": 36859.0, + "40": 37400.0, + "41": 36473.0, + "42": 36460.0, + "43": 37019.0, + "44": 36370.0, + "45": 39844.0, + "46": 38478.0, + "47": 38793.0, + "48": 38437.0, + "49": 42666.0, + "50": 39545.0 } }, "mem-allocated-bytes": { @@ -118,56 +118,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 1027089408.0, - "2": 1027091456.0, - "3": 1027087360.0, - "4": 1027088384.0, - "5": 1027091456.0, - "6": 1027091456.0, - "7": 1027088896.0, - "8": 1027092480.0, + "1": 1027088896.0, + "2": 1027090944.0, + "3": 1027086848.0, + "4": 1027087360.0, + "5": 1027091968.0, + "6": 1027089920.0, + "7": 1027088384.0, + "8": 1027091968.0, "9": 1027091968.0, - "10": 1027089408.0, + "10": 1027089920.0, "11": 1027089920.0, - "12": 1027092480.0, + "12": 1027091456.0, "13": 1027090944.0, - "14": 1027092480.0, - "15": 1027090432.0, - "16": 1027088384.0, - "17": 1027089408.0, - "18": 1027090944.0, - "19": 1027088384.0, - "20": 1027090432.0, - "21": 1027092480.0, + "14": 1027091456.0, + "15": 1027089920.0, + "16": 1027088896.0, + "17": 1027087872.0, + "18": 1027090432.0, + "19": 1027088896.0, + "20": 1027089408.0, + "21": 1027091456.0, "22": 1027089920.0, - "23": 1027093504.0, + "23": 1027094016.0, "24": 1027092480.0, "25": 1027089408.0, - "26": 1027090944.0, - "27": 1027087360.0, - "28": 1027090432.0, - "29": 1027090432.0, - "30": 1027089920.0, - "31": 1027089408.0, - "32": 1027093504.0, - "33": 1027094016.0, - "34": 1027093504.0, - "35": 1027085824.0, - "36": 1027087872.0, + "26": 1027089408.0, + "27": 1027086848.0, + "28": 1027090944.0, + "29": 1027089920.0, + "30": 1027090432.0, + "31": 1027090432.0, + "32": 1027092992.0, + "33": 1027092480.0, + "34": 1027091968.0, + "35": 1027085312.0, + "36": 1027086336.0, "37": 1027088896.0, "38": 1027089920.0, - "39": 1027088384.0, - "40": 1027091968.0, + "39": 1027087360.0, + "40": 1027090944.0, "41": 1027088384.0, - "42": 1027089408.0, + "42": 1027087872.0, "43": 1027087872.0, - "44": 1027091456.0, + "44": 1027088384.0, "45": 1027090432.0, "46": 1027086336.0, - "47": 1027088384.0, - "48": 1027087360.0, - "49": 1027087360.0, - "50": 1027089920.0 + "47": 1027086848.0, + "48": 1027086848.0, + "49": 1027086336.0, + "50": 1027089408.0 } }, "mem-max-allocated-bytes": { @@ -175,113 +175,112 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 3058326528.0, - "2": 3298517504.0, - "3": 3298517504.0, - "4": 3298517504.0, - "5": 3300747776.0, - "6": 3300747776.0, - "7": 3300747776.0, - "8": 3300747776.0, - "9": 3300747776.0, - "10": 3300747776.0, - "11": 3300747776.0, - "12": 3300747776.0, - "13": 3300747776.0, - "14": 3300747776.0, - "15": 3300747776.0, - "16": 3300747776.0, - "17": 3300747776.0, - "18": 3300747776.0, - "19": 3300747776.0, - "20": 3300747776.0, - "21": 3300747776.0, - "22": 3300747776.0, - "23": 3300747776.0, - "24": 3300747776.0, - "25": 3300747776.0, - "26": 3300747776.0, - "27": 3300747776.0, - "28": 3300747776.0, - "29": 3300747776.0, - "30": 3300747776.0, - "31": 3300747776.0, - "32": 3300747776.0, - "33": 3300747776.0, - "34": 3300872192.0, - "35": 3300872192.0, - "36": 3300872192.0, - "37": 3300872192.0, - "38": 3300872192.0, - "39": 3300872192.0, - "40": 3300872192.0, - "41": 3300872192.0, - "42": 3300872192.0, - "43": 3300872192.0, - "44": 3300872192.0, - "45": 3300872192.0, - "46": 3300872192.0, - "47": 3300872192.0, - "48": 3300872192.0, - "49": 3300872192.0, - "50": 3300872192.0 + "1": 3059096576.0, + "2": 3298776064.0, + "3": 3298776064.0, + "4": 3298776064.0, + "5": 3298776064.0, + "6": 3298776064.0, + "7": 3298776064.0, + "8": 3299136512.0, + "9": 3299136512.0, + "10": 3299136512.0, + "11": 3299136512.0, + "12": 3299136512.0, + "13": 3299320320.0, + "14": 3299397120.0, + "15": 3299397120.0, + "16": 3299397120.0, + "17": 3299397120.0, + "18": 3299397120.0, + "19": 3299397120.0, + "20": 3299397120.0, + "21": 3299397120.0, + "22": 3299397120.0, + "23": 3300247552.0, + "24": 3300247552.0, + "25": 3300247552.0, + "26": 3300247552.0, + "27": 3300247552.0, + "28": 3300247552.0, + "29": 3300247552.0, + "30": 3300247552.0, + "31": 3300247552.0, + "32": 3300247552.0, + "33": 3300247552.0, + "34": 3300554752.0, + "35": 3300554752.0, + "36": 3300554752.0, + "37": 3300554752.0, + "38": 3300554752.0, + "39": 3300554752.0, + "40": 3300554752.0, + "41": 3300554752.0, + "42": 3300554752.0, + "43": 3300554752.0, + "44": 3300554752.0, + "45": 3300554752.0, + "46": 3300554752.0, + "47": 3300554752.0, + "48": 3300554752.0, + "49": 3300554752.0, + "50": 3300554752.0 } }, "iteration-time": { - "start_step": 1, + "start_step": 2, "end_step": 50, "step_interval": 1, "values": { - "1": "nan", - "2": 7.27219, - "3": 0.24679, - "4": 0.22635, - "5": 0.23051, - "6": 0.22263, - "7": 0.21818, - "8": 0.21355, - "9": 0.21356, - "10": 0.21043, - "11": 0.21544, - "12": 0.21111, - "13": 0.21015, - "14": 0.21431, - "15": 0.21165, - "16": 0.21367, - "17": 0.21668, - "18": 0.2084, - "19": 0.20834, - "20": 0.20701, - "21": 0.21147, - "22": 0.20775, - "23": 0.2219, - "24": 0.21061, - "25": 0.20661, - "26": 0.21028, - "27": 0.2129, - "28": 0.20786, - "29": 0.20797, - "30": 0.20789, - "31": 0.20896, - "32": 0.20624, - "33": 0.20688, - "34": 0.20637, - "35": 0.20779, - "36": 0.20898, - "37": 0.20801, - "38": 0.2083, - "39": 0.20824, - "40": 0.20749, - "41": 0.20582, - "42": 0.20712, - "43": 0.21062, - "44": 0.21109, - "45": 0.21193, - "46": 0.20663, - "47": 0.21151, - "48": 0.20703, - "49": 0.21392, - "50": 0.21062 + "2": 5.16303, + "3": 0.31975, + "4": 0.32114, + "5": 0.32174, + "6": 0.31434, + "7": 0.30635, + "8": 0.30462, + "9": 0.30412, + "10": 0.30505, + "11": 0.31288, + "12": 0.29834, + "13": 0.30086, + "14": 0.29552, + "15": 0.28913, + "16": 0.29571, + "17": 0.29148, + "18": 0.2864, + "19": 0.28774, + "20": 0.28853, + "21": 0.2911, + "22": 0.29212, + "23": 0.2979, + "24": 0.29998, + "25": 0.29745, + "26": 0.29211, + "27": 0.30024, + "28": 0.29179, + "29": 0.29311, + "30": 0.29515, + "31": 0.29371, + "32": 0.29669, + "33": 0.29283, + "34": 0.29194, + "35": 0.29361, + "36": 0.2973, + "37": 0.29273, + "38": 0.29215, + "39": 0.29439, + "40": 0.295, + "41": 0.28702, + "42": 0.29175, + "43": 0.28749, + "44": 0.29187, + "45": 1.10566, + "46": 0.28901, + "47": 0.2914, + "48": 0.30221, + "49": 0.30073, + "50": 0.29095 } } -} \ No newline at end of file +} diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM/model_config.yaml index 8ced6e37a52..9353bea2933 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM/model_config.yaml @@ -64,3 +64,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM_1node/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM_1node/model_config.yaml index 8ced6e37a52..9353bea2933 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_overlap_grad_reduce_param_gather_groupedGEMM_1node/model_config.yaml @@ -64,3 +64,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_top2router/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_top2router/model_config.yaml index 097beac2085..bfaea1e6d5c 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_top2router/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts2parallel_top2router/model_config.yaml @@ -61,3 +61,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts_etp1_ep4/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts_etp1_ep4/model_config.yaml index 8ae6dc79fe2..665792146e6 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts_etp1_ep4/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_8experts_etp1_ep4/model_config.yaml @@ -64,3 +64,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4/golden_values_dev_dgx_h100.json index 06a1d993f8b..8bcf15522c7 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4/golden_values_dev_dgx_h100.json +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4/golden_values_dev_dgx_h100.json @@ -6,54 +6,54 @@ "values": { "1": 10.90791, "2": 10.90713, - "3": 10.91668, - "4": 10.90899, - "5": 10.91483, - "6": 10.89524, - "7": 10.90676, - "8": 10.90896, - "9": 10.90939, - "10": 10.91077, - "11": 10.901, - "12": 10.89922, - "13": 10.88807, - "14": 10.88197, - "15": 10.87251, - "16": 10.85287, - "17": 10.85704, - "18": 10.84826, - "19": 10.85507, - "20": 10.77687, - "21": 10.76073, - "22": 10.7605, - "23": 10.74325, - "24": 10.70838, - "25": 10.70981, - "26": 10.69235, - "27": 10.66868, - "28": 10.60599, - "29": 10.57223, - "30": 10.54151, - "31": 10.53199, - "32": 10.51634, - "33": 10.481, - "34": 10.44913, - "35": 10.44632, - "36": 10.42066, - "37": 10.40067, - "38": 10.40454, - "39": 10.36981, - "40": 10.35244, - "41": 10.33008, - "42": 10.31128, - "43": 10.29795, - "44": 10.27171, - "45": 10.28363, - "46": 10.24114, - "47": 10.23434, - "48": 10.19197, - "49": 10.19498, - "50": 10.19073 + "3": 10.91656, + "4": 10.90856, + "5": 10.91486, + "6": 10.8955, + "7": 10.90682, + "8": 10.90938, + "9": 10.90906, + "10": 10.91044, + "11": 10.90161, + "12": 10.9002, + "13": 10.88758, + "14": 10.88178, + "15": 10.87374, + "16": 10.85236, + "17": 10.85613, + "18": 10.84761, + "19": 10.85533, + "20": 10.77576, + "21": 10.76185, + "22": 10.75979, + "23": 10.74357, + "24": 10.7085, + "25": 10.70954, + "26": 10.69236, + "27": 10.66744, + "28": 10.60538, + "29": 10.57197, + "30": 10.541, + "31": 10.53121, + "32": 10.51574, + "33": 10.48005, + "34": 10.44884, + "35": 10.44593, + "36": 10.41946, + "37": 10.39946, + "38": 10.40313, + "39": 10.36823, + "40": 10.35085, + "41": 10.3285, + "42": 10.30918, + "43": 10.29587, + "44": 10.26912, + "45": 10.28118, + "46": 10.2386, + "47": 10.23166, + "48": 10.18908, + "49": 10.19252, + "50": 10.1881 } }, "num-zeros": { @@ -61,56 +61,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 16592.0, - "2": 16504.0, - "3": 16575.0, - "4": 16376.0, - "5": 16183.0, - "6": 16014.0, - "7": 16818.0, - "8": 15883.0, - "9": 16693.0, - "10": 16580.0, - "11": 16233.0, - "12": 16030.0, - "13": 16714.0, - "14": 16720.0, - "15": 15895.0, - "16": 16103.0, - "17": 16623.0, - "18": 16575.0, - "19": 16738.0, - "20": 16598.0, - "21": 16235.0, - "22": 16514.0, - "23": 16299.0, - "24": 16586.0, - "25": 15770.0, - "26": 16540.0, - "27": 16550.0, - "28": 16182.0, - "29": 16694.0, - "30": 16454.0, - "31": 17100.0, - "32": 17390.0, - "33": 17062.0, - "34": 17233.0, - "35": 17703.0, - "36": 17351.0, - "37": 18011.0, - "38": 17621.0, - "39": 18243.0, - "40": 18971.0, - "41": 18220.0, - "42": 17966.0, - "43": 18752.0, - "44": 18809.0, - "45": 20890.0, - "46": 19846.0, - "47": 19418.0, - "48": 20136.0, - "49": 22380.0, - "50": 20145.0 + "1": 32481.0, + "2": 32205.0, + "3": 31782.0, + "4": 32082.0, + "5": 31672.0, + "6": 30901.0, + "7": 32296.0, + "8": 30851.0, + "9": 32347.0, + "10": 32485.0, + "11": 31812.0, + "12": 31039.0, + "13": 32298.0, + "14": 32795.0, + "15": 31592.0, + "16": 30976.0, + "17": 32064.0, + "18": 32220.0, + "19": 32630.0, + "20": 32593.0, + "21": 31945.0, + "22": 32124.0, + "23": 32052.0, + "24": 33318.0, + "25": 31411.0, + "26": 32629.0, + "27": 32773.0, + "28": 32484.0, + "29": 32771.0, + "30": 32994.0, + "31": 34132.0, + "32": 34806.0, + "33": 33924.0, + "34": 34187.0, + "35": 35432.0, + "36": 35117.0, + "37": 35331.0, + "38": 35038.0, + "39": 36823.0, + "40": 38166.0, + "41": 36109.0, + "42": 35997.0, + "43": 37458.0, + "44": 37012.0, + "45": 40824.0, + "46": 38797.0, + "47": 38754.0, + "48": 39579.0, + "49": 43844.0, + "50": 39384.0 } }, "mem-allocated-bytes": { @@ -118,56 +118,56 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 1562452480.0, - "2": 1560886272.0, - "3": 1560936960.0, - "4": 1560987648.0, - "5": 1560990208.0, - "6": 1560987648.0, - "7": 1561697792.0, - "8": 1561797632.0, - "9": 1560936960.0, - "10": 1560987648.0, - "11": 1562030592.0, - "12": 1561056256.0, - "13": 1561038336.0, - "14": 1561974784.0, - "15": 1561152000.0, - "16": 1561089024.0, - "17": 1561038336.0, - "18": 1562183680.0, - "19": 1562019328.0, - "20": 1561089024.0, - "21": 1561960960.0, - "22": 1561631744.0, - "23": 1561836032.0, - "24": 1561089024.0, - "25": 1561139712.0, - "26": 1561089024.0, - "27": 1561139712.0, - "28": 1561190400.0, - "29": 1561772544.0, - "30": 1561604096.0, - "31": 1562013184.0, - "32": 1561406464.0, - "33": 1561139712.0, - "34": 1561755648.0, - "35": 1561589248.0, - "36": 1561190400.0, - "37": 1561662976.0, - "38": 1561190400.0, - "39": 1561241088.0, - "40": 1563022336.0, - "41": 1564137984.0, - "42": 1561291776.0, - "43": 1561952768.0, - "44": 1561291776.0, - "45": 1561715200.0, - "46": 1561291776.0, - "47": 1561342464.0, - "48": 1561291776.0, - "49": 1561342464.0, - "50": 1561291776.0 + "1": 1564053504.0, + "2": 1563238912.0, + "3": 1562403328.0, + "4": 1562403328.0, + "5": 1562403328.0, + "6": 1562403328.0, + "7": 1562403328.0, + "8": 1562513920.0, + "9": 1563245056.0, + "10": 1562403328.0, + "11": 1562484224.0, + "12": 1562403328.0, + "13": 1562433024.0, + "14": 1562863104.0, + "15": 1563126272.0, + "16": 1563629056.0, + "17": 1563629056.0, + "18": 1564133888.0, + "19": 1562616320.0, + "20": 1562403328.0, + "21": 1562403328.0, + "22": 1562895872.0, + "23": 1563697664.0, + "24": 1562403328.0, + "25": 1562640896.0, + "26": 1562403328.0, + "27": 1562403328.0, + "28": 1562403328.0, + "29": 1562691072.0, + "30": 1563299328.0, + "31": 1562403328.0, + "32": 1562403328.0, + "33": 1562915328.0, + "34": 1563009536.0, + "35": 1562771968.0, + "36": 1562403328.0, + "37": 1563407872.0, + "38": 1563094528.0, + "39": 1562403328.0, + "40": 1562403328.0, + "41": 1563132416.0, + "42": 1562603008.0, + "43": 1562563072.0, + "44": 1562528256.0, + "45": 1562403328.0, + "46": 1562589696.0, + "47": 1562403328.0, + "48": 1562632704.0, + "49": 1562403328.0, + "50": 1563123200.0 } }, "mem-max-allocated-bytes": { @@ -175,113 +175,112 @@ "end_step": 50, "step_interval": 1, "values": { - "1": 3479876608.0, - "2": 4042335744.0, - "3": 4047595520.0, - "4": 4053274112.0, - "5": 4053274112.0, - "6": 4053274112.0, - "7": 4057433088.0, - "8": 4061642240.0, - "9": 4061642240.0, - "10": 4065529344.0, - "11": 4065529344.0, - "12": 4065529344.0, - "13": 4065529344.0, - "14": 4065529344.0, - "15": 4065529344.0, - "16": 4065529344.0, - "17": 4065529344.0, - "18": 4065529344.0, - "19": 4065529344.0, - "20": 4065529344.0, - "21": 4065529344.0, - "22": 4065529344.0, - "23": 4065529344.0, - "24": 4065529344.0, - "25": 4065529344.0, - "26": 4065529344.0, - "27": 4065529344.0, - "28": 4065529344.0, - "29": 4065529344.0, - "30": 4065529344.0, - "31": 4065529344.0, - "32": 4065529344.0, - "33": 4065529344.0, - "34": 4065529344.0, - "35": 4065529344.0, - "36": 4065529344.0, - "37": 4065529344.0, - "38": 4065529344.0, - "39": 4065529344.0, - "40": 4065529344.0, - "41": 4065529344.0, - "42": 4065529344.0, - "43": 4065529344.0, - "44": 4065529344.0, - "45": 4065529344.0, - "46": 4065529344.0, - "47": 4065529344.0, - "48": 4065529344.0, - "49": 4065529344.0, - "50": 4065529344.0 + "1": 3481845248.0, + "2": 4045780992.0, + "3": 4052392960.0, + "4": 4057906176.0, + "5": 4057906176.0, + "6": 4057906176.0, + "7": 4057906176.0, + "8": 4057906176.0, + "9": 4060597248.0, + "10": 4065854464.0, + "11": 4065854464.0, + "12": 4065854464.0, + "13": 4065854464.0, + "14": 4065854464.0, + "15": 4065854464.0, + "16": 4065854464.0, + "17": 4065854464.0, + "18": 4065854464.0, + "19": 4065854464.0, + "20": 4065854464.0, + "21": 4065854464.0, + "22": 4065854464.0, + "23": 4065854464.0, + "24": 4065854464.0, + "25": 4065854464.0, + "26": 4065854464.0, + "27": 4065854464.0, + "28": 4065854464.0, + "29": 4065854464.0, + "30": 4065854464.0, + "31": 4065854464.0, + "32": 4065854464.0, + "33": 4065854464.0, + "34": 4065854464.0, + "35": 4065854464.0, + "36": 4065854464.0, + "37": 4065854464.0, + "38": 4065854464.0, + "39": 4065854464.0, + "40": 4065854464.0, + "41": 4065854464.0, + "42": 4065854464.0, + "43": 4065854464.0, + "44": 4065854464.0, + "45": 4065854464.0, + "46": 4065854464.0, + "47": 4065854464.0, + "48": 4065854464.0, + "49": 4065854464.0, + "50": 4065854464.0 } }, "iteration-time": { - "start_step": 1, + "start_step": 2, "end_step": 50, "step_interval": 1, "values": { - "1": "nan", - "2": 11.44286, - "3": 0.39137, - "4": 0.33071, - "5": 0.32257, - "6": 0.32404, - "7": 0.30595, - "8": 0.29297, - "9": 0.29395, - "10": 0.27912, - "11": 0.30251, - "12": 0.28669, - "13": 0.28455, - "14": 0.28124, - "15": 0.2876, - "16": 0.27705, - "17": 0.28277, - "18": 0.28818, - "19": 0.29518, - "20": 0.28783, - "21": 0.28453, - "22": 0.28955, - "23": 0.27766, - "24": 0.278, - "25": 0.28149, - "26": 0.29603, - "27": 0.27934, - "28": 0.29048, - "29": 0.29607, - "30": 0.28981, - "31": 0.32857, - "32": 0.29071, - "33": 0.29613, - "34": 0.2968, - "35": 0.30616, - "36": 0.30069, - "37": 0.29431, - "38": 0.29876, - "39": 0.30582, - "40": 0.28349, - "41": 0.28535, - "42": 0.28254, - "43": 0.2788, - "44": 0.27508, - "45": 0.27863, - "46": 0.27541, - "47": 0.27561, - "48": 0.27969, - "49": 0.27721, - "50": 0.27313 + "2": 7.65129, + "3": 0.49113, + "4": 0.46732, + "5": 0.45015, + "6": 0.44098, + "7": 0.43279, + "8": 0.43532, + "9": 0.41644, + "10": 0.41447, + "11": 0.417, + "12": 1.27039, + "13": 0.42857, + "14": 0.42043, + "15": 0.43429, + "16": 0.42646, + "17": 0.41839, + "18": 0.41875, + "19": 0.42078, + "20": 0.4152, + "21": 0.41839, + "22": 0.42007, + "23": 0.40978, + "24": 0.4028, + "25": 0.40842, + "26": 0.41505, + "27": 1.15357, + "28": 0.43831, + "29": 1.14215, + "30": 0.42585, + "31": 0.42393, + "32": 0.42386, + "33": 0.41322, + "34": 0.42071, + "35": 0.4168, + "36": 1.17628, + "37": 0.42372, + "38": 0.42557, + "39": 1.14457, + "40": 0.4147, + "41": 0.41313, + "42": 0.41232, + "43": 0.41219, + "44": 0.41084, + "45": 0.40381, + "46": 0.40844, + "47": 0.40717, + "48": 1.20976, + "49": 0.42107, + "50": 0.42035 } } -} \ No newline at end of file +} diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4/model_config.yaml index 5423ee39527..d7711454e7f 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4/model_config.yaml @@ -66,3 +66,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4_1node/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4_1node/model_config.yaml index 5423ee39527..d7711454e7f 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp1_te_a2a_ovlp_8experts_etp1_ep4_1node/model_config.yaml @@ -66,3 +66,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_memory_speed/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_memory_speed/model_config.yaml index 8f108eabdac..9f046fa10ff 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_memory_speed/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_memory_speed/model_config.yaml @@ -133,3 +133,4 @@ METRICS: - "mem-allocated-bytes" - "mem-max-allocated-bytes" - "mtp_1 loss" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_mtp_resume_torch_dist_fp8/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_mtp_resume_torch_dist_fp8/model_config.yaml index 3f1ce0e8f16..7976caff8ec 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_mtp_resume_torch_dist_fp8/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_mtp_resume_torch_dist_fp8/model_config.yaml @@ -135,3 +135,4 @@ METRICS: - "mem-allocated-bytes" - "mem-max-allocated-bytes" - "mtp_1 loss" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_resume_torch_dist_attn_cudagraph/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_resume_torch_dist_attn_cudagraph/model_config.yaml index 64cdacd6076..c534cf14ec3 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_resume_torch_dist_attn_cudagraph/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_resume_torch_dist_attn_cudagraph/model_config.yaml @@ -137,3 +137,4 @@ METRICS: - "num-zeros" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_selective_recompute_experimental/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_selective_recompute_experimental/model_config.yaml index 172fac96f6f..b107e806902 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_selective_recompute_experimental/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_pp2_ep4_etp1_selective_recompute_experimental/model_config.yaml @@ -136,3 +136,4 @@ METRICS: - "mem-allocated-bytes" - "mem-max-allocated-bytes" - "mtp_1 loss" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_zp_z3_resume_torch_dist_te_8experts2parallel_top2router/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_zp_z3_resume_torch_dist_te_8experts2parallel_top2router/model_config.yaml index 15f971e9ff3..4e060b589c2 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_zp_z3_resume_torch_dist_te_8experts2parallel_top2router/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_te_tp2_zp_z3_resume_torch_dist_te_8experts2parallel_top2router/model_config.yaml @@ -63,3 +63,4 @@ MODEL_ARGS: --bf16: true --no-bias-gelu-fusion: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_cp2_pp2_ep2_te_4experts2parallel/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_cp2_pp2_ep2_te_4experts2parallel/model_config.yaml index 901cb22f005..a2cff5cdab5 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_cp2_pp2_ep2_te_4experts2parallel/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_cp2_pp2_ep2_te_4experts2parallel/model_config.yaml @@ -60,3 +60,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_cp2_pp2_ep2_te_4experts2parallel_dp_last/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_cp2_pp2_ep2_te_4experts2parallel_dp_last/model_config.yaml index 61d25aeb356..f6a8b3f365c 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_cp2_pp2_ep2_te_4experts2parallel_dp_last/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_cp2_pp2_ep2_te_4experts2parallel_dp_last/model_config.yaml @@ -60,3 +60,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_etp2_te_4experts2parallel/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_etp2_te_4experts2parallel/model_config.yaml index b03dd7fe023..0a3087e9bcb 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_etp2_te_4experts2parallel/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_etp2_te_4experts2parallel/model_config.yaml @@ -59,3 +59,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_etp2_te_4experts2parallel_dp_last/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_etp2_te_4experts2parallel_dp_last/model_config.yaml index c48348735b8..261074e80a1 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_etp2_te_4experts2parallel_dp_last/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_etp2_te_4experts2parallel_dp_last/model_config.yaml @@ -60,3 +60,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_resume_torch_dist_te_4experts2parallel/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_resume_torch_dist_te_4experts2parallel/model_config.yaml index 192f7d09101..b32a30d946c 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_resume_torch_dist_te_4experts2parallel/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_resume_torch_dist_te_4experts2parallel/model_config.yaml @@ -58,3 +58,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_te_4experts2parallel/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_te_4experts2parallel/model_config.yaml index f4b020370ff..2ec13052640 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_te_4experts2parallel/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_ep2_te_4experts2parallel/model_config.yaml @@ -59,3 +59,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_resume_torch_dist_te_2experts/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_resume_torch_dist_te_2experts/model_config.yaml index cf0f282e5b1..efe1fa34a8f 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_resume_torch_dist_te_2experts/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_mcore_tp2_pp2_resume_torch_dist_te_2experts/model_config.yaml @@ -58,3 +58,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon/model_config.yaml index 028074bd34f..bc3bd5484f1 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon/model_config.yaml @@ -67,3 +67,4 @@ MODEL_ARGS: --use-persistent-ckpt-worker: true --use-distributed-optimizer: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_1node/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_1node/model_config.yaml index 40ac94eed9b..d184ee46673 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_1node/model_config.yaml @@ -67,3 +67,4 @@ MODEL_ARGS: --use-persistent-ckpt-worker: true --use-distributed-optimizer: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_param_layout/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_param_layout/model_config.yaml index 028074bd34f..bc3bd5484f1 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_param_layout/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_param_layout/model_config.yaml @@ -67,3 +67,4 @@ MODEL_ARGS: --use-persistent-ckpt-worker: true --use-distributed-optimizer: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_param_layout_1node/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_param_layout_1node/model_config.yaml index 40ac94eed9b..d184ee46673 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_param_layout_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_muon_param_layout_1node/model_config.yaml @@ -67,3 +67,4 @@ MODEL_ARGS: --use-persistent-ckpt-worker: true --use-distributed-optimizer: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_optimizer/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_optimizer/model_config.yaml index f69a44638d6..317bd508caf 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_optimizer/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_optimizer/model_config.yaml @@ -65,3 +65,4 @@ MODEL_ARGS: --async-strategy: mcore --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_optimizer_1node/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_optimizer_1node/model_config.yaml index 85b4ea629df..1aabaee0efa 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_optimizer_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_dist_optimizer_1node/model_config.yaml @@ -65,3 +65,4 @@ MODEL_ARGS: --async-strategy: mcore --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_muon/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_muon/model_config.yaml index 6b14d162d9f..d5ee0d0fe8c 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_muon/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_muon/model_config.yaml @@ -68,3 +68,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_muon_1node/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_muon_1node/model_config.yaml index e1653ce7ed4..736b3c52f9f 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_muon_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_ep8_resume_torch_dist_muon_1node/model_config.yaml @@ -68,3 +68,4 @@ MODEL_ARGS: --async-save: true --use-persistent-ckpt-worker: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp2_pp2_ep4_etp1_fine_grained_offloading/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp2_pp2_ep4_etp1_fine_grained_offloading/model_config.yaml index 845f0990460..fcab06c39fc 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp2_pp2_ep4_etp1_fine_grained_offloading/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp2_pp2_ep4_etp1_fine_grained_offloading/model_config.yaml @@ -140,3 +140,4 @@ METRICS: - "mem-allocated-bytes" - "mem-max-allocated-bytes" - "mtp_1 loss" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp2_pp2_ep4_etp1_no_mtp_no_a2a_ovlp_fine_grained_offloading/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp2_pp2_ep4_etp1_no_mtp_no_a2a_ovlp_fine_grained_offloading/model_config.yaml index d16c56b2264..a639727bc2c 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp2_pp2_ep4_etp1_no_mtp_no_a2a_ovlp_fine_grained_offloading/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp2_pp2_ep4_etp1_no_mtp_no_a2a_ovlp_fine_grained_offloading/model_config.yaml @@ -134,3 +134,4 @@ METRICS: - "lm loss" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_resume_torch_dist_dist_optimizer/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_resume_torch_dist_dist_optimizer/golden_values_dev_dgx_h100.json index 635fcea4a97..455d3cbe1c4 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_resume_torch_dist_dist_optimizer/golden_values_dev_dgx_h100.json +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_resume_torch_dist_dist_optimizer/golden_values_dev_dgx_h100.json @@ -7,103 +7,103 @@ "1": 10.92486, "2": 10.91069, "3": 10.91839, - "4": 10.91715, - "5": 10.90503, - "6": 10.90196, - "7": 10.89732, - "8": 10.91344, - "9": 10.91625, - "10": 10.91028, - "11": 10.90163, - "12": 10.89726, - "13": 10.8879, - "14": 10.89483, - "15": 10.87506, - "16": 10.87056, - "17": 10.86912, - "18": 10.85168, - "19": 10.87022, - "20": 10.78801, - "21": 10.77233, - "22": 10.76715, - "23": 10.75842, - "24": 10.71926, - "25": 10.71997, - "26": 10.71229, - "27": 10.68551, - "28": 10.61305, - "29": 10.58637, - "30": 10.56568, - "31": 10.55773, - "32": 10.54888, - "33": 10.50977, - "34": 10.48172, - "35": 10.47011, - "36": 10.45293, - "37": 10.42772, - "38": 10.43271, - "39": 10.40299, - "40": 10.3775, - "41": 10.36875, - "42": 10.33113, - "43": 10.31542, - "44": 10.29012, - "45": 10.30282, - "46": 10.2657, - "47": 10.25567, - "48": 10.20709, - "49": 10.21058, - "50": 10.21068, + "4": 10.9172, + "5": 10.90494, + "6": 10.90213, + "7": 10.8971, + "8": 10.91319, + "9": 10.9167, + "10": 10.9103, + "11": 10.90137, + "12": 10.8968, + "13": 10.88802, + "14": 10.89532, + "15": 10.87553, + "16": 10.87001, + "17": 10.86926, + "18": 10.8521, + "19": 10.86958, + "20": 10.78798, + "21": 10.77234, + "22": 10.7674, + "23": 10.75859, + "24": 10.7193, + "25": 10.72025, + "26": 10.71224, + "27": 10.68505, + "28": 10.61329, + "29": 10.5867, + "30": 10.56566, + "31": 10.55778, + "32": 10.54894, + "33": 10.5097, + "34": 10.48124, + "35": 10.46985, + "36": 10.45295, + "37": 10.42776, + "38": 10.43249, + "39": 10.40289, + "40": 10.37722, + "41": 10.36871, + "42": 10.33138, + "43": 10.31516, + "44": 10.29024, + "45": 10.3027, + "46": 10.26552, + "47": 10.25565, + "48": 10.20692, + "49": 10.21051, + "50": 10.21041, "51": 10.21195, - "52": 10.16251, - "53": 10.16325, - "54": 10.13402, - "55": 10.10872, - "56": 10.13453, - "57": 10.13277, - "58": 10.12413, - "59": 10.06521, - "60": 10.09524, - "61": 10.04762, - "62": 10.01553, - "63": 10.08297, - "64": 10.03279, - "65": 9.99846, - "66": 10.03919, - "67": 10.01295, - "68": 9.97759, - "69": 9.99341, - "70": 9.97104, - "71": 9.9982, - "72": 9.97568, - "73": 9.95991, - "74": 9.95299, - "75": 9.9144, - "76": 9.95006, - "77": 9.94205, - "78": 9.89904, - "79": 9.89709, - "80": 9.91042, - "81": 9.93365, - "82": 9.88344, - "83": 9.83983, - "84": 9.7821, - "85": 9.76275, - "86": 9.87789, - "87": 9.9007, - "88": 9.87402, - "89": 9.82463, - "90": 9.81377, - "91": 9.8198, - "92": 9.81606, - "93": 9.74349, - "94": 9.82172, - "95": 9.81227, - "96": 9.79491, - "97": 9.74649, - "98": 9.7688, - "99": 9.81824, - "100": 9.70741 + "52": 10.16255, + "53": 10.1631, + "54": 10.13394, + "55": 10.10862, + "56": 10.13462, + "57": 10.13253, + "58": 10.12393, + "59": 10.06506, + "60": 10.09504, + "61": 10.04746, + "62": 10.01501, + "63": 10.08262, + "64": 10.03252, + "65": 9.99817, + "66": 10.03863, + "67": 10.01262, + "68": 9.97709, + "69": 9.99273, + "70": 9.97043, + "71": 9.99761, + "72": 9.97478, + "73": 9.95903, + "74": 9.95198, + "75": 9.91332, + "76": 9.9493, + "77": 9.94113, + "78": 9.89802, + "79": 9.89619, + "80": 9.90937, + "81": 9.93263, + "82": 9.88265, + "83": 9.83871, + "84": 9.78108, + "85": 9.76154, + "86": 9.87689, + "87": 9.89981, + "88": 9.87312, + "89": 9.82362, + "90": 9.81273, + "91": 9.81873, + "92": 9.81493, + "93": 9.74223, + "94": 9.82036, + "95": 9.81099, + "96": 9.79374, + "97": 9.74486, + "98": 9.76728, + "99": 9.81697, + "100": 9.70593 } }, "num-zeros": { @@ -111,106 +111,106 @@ "end_step": 100, "step_interval": 1, "values": { - "1": 2429.0, - "2": 2591.0, - "3": 2604.0, - "4": 2665.0, - "5": 2610.0, - "6": 2487.0, - "7": 2589.0, - "8": 2535.0, - "9": 2677.0, - "10": 2474.0, - "11": 2585.0, - "12": 2605.0, - "13": 2530.0, - "14": 2651.0, - "15": 2509.0, - "16": 2506.0, - "17": 2646.0, - "18": 2601.0, - "19": 2552.0, - "20": 2518.0, - "21": 2539.0, - "22": 2594.0, - "23": 2531.0, - "24": 2604.0, - "25": 2474.0, - "26": 2505.0, - "27": 2647.0, - "28": 2551.0, - "29": 2735.0, - "30": 2709.0, - "31": 2746.0, - "32": 2729.0, - "33": 2672.0, - "34": 2741.0, - "35": 2722.0, - "36": 2761.0, - "37": 2860.0, - "38": 2827.0, - "39": 3030.0, - "40": 3060.0, - "41": 3129.0, - "42": 2813.0, - "43": 3059.0, - "44": 3088.0, - "45": 3301.0, - "46": 3239.0, - "47": 3241.0, - "48": 3278.0, - "49": 3483.0, - "50": 3340.0, - "51": 3328.0, - "52": 3384.0, - "53": 3253.0, - "54": 3558.0, - "55": 3290.0, - "56": 3548.0, - "57": 2934.0, - "58": 3989.0, - "59": 3538.0, - "60": 3638.0, - "61": 3440.0, - "62": 3763.0, - "63": 3857.0, - "64": 4201.0, - "65": 3318.0, - "66": 3743.0, - "67": 4019.0, - "68": 3853.0, - "69": 3501.0, - "70": 3812.0, - "71": 3749.0, - "72": 3597.0, - "73": 4178.0, - "74": 3676.0, - "75": 3677.0, - "76": 4080.0, - "77": 4149.0, - "78": 4236.0, - "79": 7606.0, - "80": 33298.0, - "81": 8170.0, - "82": 528771.0, - "83": 3458.0, - "84": 31908.0, - "85": 528749.0, - "86": 529194.0, - "87": 61926.0, - "88": 529038.0, - "89": 529251.0, - "90": 28836.0, - "91": 528719.0, - "92": 529072.0, - "93": 1053703.0, - "94": 529234.0, - "95": 553148.0, - "96": 560606.0, - "97": 529810.0, - "98": 529332.0, - "99": 529265.0, - "100": 529071.0 + "1": 6308.0, + "2": 6526.0, + "3": 6460.0, + "4": 6587.0, + "5": 6499.0, + "6": 6370.0, + "7": 6708.0, + "8": 6525.0, + "9": 6677.0, + "10": 6625.0, + "11": 6727.0, + "12": 6349.0, + "13": 6508.0, + "14": 6824.0, + "15": 6084.0, + "16": 6481.0, + "17": 6610.0, + "18": 6402.0, + "19": 6412.0, + "20": 6120.0, + "21": 6510.0, + "22": 6553.0, + "23": 6598.0, + "24": 6752.0, + "25": 6592.0, + "26": 6412.0, + "27": 6775.0, + "28": 6714.0, + "29": 7049.0, + "30": 6871.0, + "31": 7154.0, + "32": 7296.0, + "33": 6998.0, + "34": 7308.0, + "35": 7361.0, + "36": 7195.0, + "37": 7698.0, + "38": 7541.0, + "39": 7777.0, + "40": 7986.0, + "41": 8348.0, + "42": 7583.0, + "43": 8268.0, + "44": 7990.0, + "45": 8716.0, + "46": 8372.0, + "47": 8571.0, + "48": 8629.0, + "49": 8993.0, + "50": 8812.0, + "51": 8718.0, + "52": 9112.0, + "53": 8324.0, + "54": 9142.0, + "55": 8346.0, + "56": 9464.0, + "57": 7897.0, + "58": 10393.0, + "59": 9474.0, + "60": 9199.0, + "61": 8898.0, + "62": 9492.0, + "63": 10239.0, + "64": 10243.0, + "65": 8621.0, + "66": 9522.0, + "67": 10217.0, + "68": 9767.0, + "69": 8996.0, + "70": 9595.0, + "71": 9956.0, + "72": 9400.0, + "73": 10200.0, + "74": 14258.0, + "75": 8955.0, + "76": 10156.0, + "77": 10151.0, + "78": 10834.0, + "79": 76326.0, + "80": 132648.0, + "81": 80347.0, + "82": 1118255.0, + "83": 66887.0, + "84": 1113524.0, + "85": 2106481.0, + "86": 2107956.0, + "87": 290100.0, + "88": 2107621.0, + "89": 2107900.0, + "90": 58643.0, + "91": 2106770.0, + "92": 2107002.0, + "93": 3155613.0, + "94": 2107794.0, + "95": 2155564.0, + "96": 2170246.0, + "97": 2108416.0, + "98": 2107390.0, + "99": 2107633.0, + "100": 2107960.0 } }, "mem-allocated-bytes": { @@ -218,106 +218,106 @@ "end_step": 100, "step_interval": 1, "values": { - "1": 628064256.0, - "2": 628065280.0, - "3": 628065280.0, - "4": 628065280.0, - "5": 628065280.0, - "6": 628065280.0, - "7": 628065280.0, - "8": 628065280.0, - "9": 628065280.0, - "10": 628065280.0, - "11": 628065280.0, - "12": 628065280.0, - "13": 628065280.0, - "14": 628065280.0, - "15": 628065280.0, - "16": 628065280.0, - "17": 628065280.0, - "18": 628065280.0, - "19": 628065280.0, - "20": 628065280.0, - "21": 628065280.0, - "22": 628065280.0, - "23": 628065280.0, - "24": 628065280.0, - "25": 628065280.0, - "26": 628065280.0, - "27": 628065280.0, - "28": 628065280.0, - "29": 628065280.0, - "30": 628065280.0, - "31": 628065280.0, - "32": 628065280.0, - "33": 628065280.0, - "34": 628065280.0, - "35": 628065280.0, - "36": 628065280.0, - "37": 628065280.0, - "38": 628065280.0, - "39": 628065280.0, - "40": 628065280.0, - "41": 628065280.0, - "42": 628065280.0, - "43": 628065280.0, - "44": 628065280.0, - "45": 628065280.0, - "46": 628065280.0, - "47": 628065280.0, - "48": 628065280.0, - "49": 628065280.0, - "50": 628065280.0, - "51": 628065280.0, - "52": 628065280.0, - "53": 628065280.0, - "54": 628065280.0, - "55": 628065280.0, - "56": 628065280.0, - "57": 628065280.0, - "58": 628065280.0, - "59": 628065280.0, - "60": 628065280.0, - "61": 628065280.0, - "62": 628065280.0, - "63": 628065280.0, - "64": 628065280.0, - "65": 628065280.0, - "66": 628065280.0, - "67": 628065280.0, - "68": 628065280.0, - "69": 628065280.0, - "70": 628065280.0, - "71": 628065280.0, - "72": 628065280.0, - "73": 628065280.0, - "74": 628065280.0, - "75": 628065280.0, - "76": 628065280.0, - "77": 628065280.0, - "78": 628065280.0, - "79": 628065280.0, - "80": 628065280.0, - "81": 628065280.0, - "82": 628065280.0, - "83": 628065280.0, - "84": 628065280.0, - "85": 628065280.0, - "86": 628065280.0, - "87": 628065280.0, - "88": 628065280.0, - "89": 628065280.0, - "90": 628065280.0, - "91": 628065280.0, - "92": 628065280.0, - "93": 628065280.0, - "94": 628065280.0, - "95": 628065280.0, - "96": 628065280.0, - "97": 628065280.0, - "98": 628065280.0, - "99": 628065280.0, - "100": 628065280.0 + "1": 628063744.0, + "2": 628064768.0, + "3": 628064768.0, + "4": 628064768.0, + "5": 628064768.0, + "6": 628064768.0, + "7": 628064768.0, + "8": 628064768.0, + "9": 628064768.0, + "10": 628064768.0, + "11": 628064768.0, + "12": 628064768.0, + "13": 628064768.0, + "14": 628064768.0, + "15": 628064768.0, + "16": 628064768.0, + "17": 628064768.0, + "18": 628064768.0, + "19": 628064768.0, + "20": 628064768.0, + "21": 628064768.0, + "22": 628064768.0, + "23": 628064768.0, + "24": 628064768.0, + "25": 628064768.0, + "26": 628064768.0, + "27": 628064768.0, + "28": 628064768.0, + "29": 628064768.0, + "30": 628064768.0, + "31": 628064768.0, + "32": 628064768.0, + "33": 628064768.0, + "34": 628064768.0, + "35": 628064768.0, + "36": 628064768.0, + "37": 628064768.0, + "38": 628064768.0, + "39": 628064768.0, + "40": 628064768.0, + "41": 628064768.0, + "42": 628064768.0, + "43": 628064768.0, + "44": 628064768.0, + "45": 628064768.0, + "46": 628064768.0, + "47": 628064768.0, + "48": 628064768.0, + "49": 628064768.0, + "50": 628064768.0, + "51": 628064768.0, + "52": 628064768.0, + "53": 628064768.0, + "54": 628064768.0, + "55": 628064768.0, + "56": 628064768.0, + "57": 628064768.0, + "58": 628064768.0, + "59": 628064768.0, + "60": 628064768.0, + "61": 628064768.0, + "62": 628064768.0, + "63": 628064768.0, + "64": 628064768.0, + "65": 628064768.0, + "66": 628064768.0, + "67": 628064768.0, + "68": 628064768.0, + "69": 628064768.0, + "70": 628064768.0, + "71": 628064768.0, + "72": 628064768.0, + "73": 628064768.0, + "74": 628064768.0, + "75": 628064768.0, + "76": 628064768.0, + "77": 628064768.0, + "78": 628064768.0, + "79": 628064768.0, + "80": 628064768.0, + "81": 628064768.0, + "82": 628064768.0, + "83": 628064768.0, + "84": 628064768.0, + "85": 628064768.0, + "86": 628064768.0, + "87": 628064768.0, + "88": 628064768.0, + "89": 628064768.0, + "90": 628064768.0, + "91": 628064768.0, + "92": 628064768.0, + "93": 628064768.0, + "94": 628064768.0, + "95": 628064768.0, + "96": 628064768.0, + "97": 628064768.0, + "98": 628064768.0, + "99": 628064768.0, + "100": 628064768.0 } }, "mem-max-allocated-bytes": { @@ -325,213 +325,212 @@ "end_step": 100, "step_interval": 1, "values": { - "1": 974434304.0, - "2": 1143423488.0, - "3": 1143897600.0, - "4": 1146645504.0, - "5": 1147860992.0, - "6": 1147860992.0, - "7": 1148556800.0, - "8": 1148556800.0, - "9": 1148556800.0, - "10": 1148556800.0, - "11": 1148556800.0, - "12": 1148556800.0, - "13": 1148556800.0, - "14": 1148556800.0, - "15": 1148556800.0, - "16": 1148556800.0, - "17": 1148556800.0, - "18": 1148556800.0, - "19": 1148556800.0, - "20": 1148556800.0, - "21": 1148556800.0, - "22": 1148556800.0, - "23": 1148556800.0, - "24": 1148556800.0, - "25": 1148556800.0, - "26": 1149693440.0, - "27": 1149693440.0, - "28": 1149693440.0, - "29": 1149693440.0, - "30": 1149693440.0, - "31": 1149693440.0, - "32": 1149693440.0, - "33": 1149693440.0, - "34": 1149693440.0, - "35": 1149693440.0, - "36": 1149693440.0, - "37": 1149693440.0, - "38": 1149693440.0, - "39": 1149693440.0, - "40": 1149693440.0, - "41": 1149693440.0, - "42": 1149693440.0, - "43": 1149693440.0, - "44": 1149693440.0, - "45": 1149693440.0, - "46": 1149693440.0, - "47": 1149693440.0, - "48": 1149693440.0, - "49": 1149693440.0, - "50": 1149693440.0, - "51": 1149693440.0, - "52": 1149693440.0, - "53": 1149693440.0, - "54": 1149693440.0, - "55": 1149693440.0, - "56": 1149693440.0, - "57": 1149693440.0, - "58": 1149693440.0, - "59": 1149693440.0, - "60": 1149693440.0, - "61": 1149693440.0, - "62": 1149693440.0, - "63": 1149693440.0, - "64": 1149693440.0, - "65": 1149693440.0, - "66": 1149693440.0, - "67": 1149693440.0, - "68": 1149693440.0, - "69": 1149693440.0, - "70": 1149693440.0, - "71": 1149693440.0, - "72": 1149693440.0, - "73": 1149693440.0, - "74": 1149693440.0, - "75": 1149693440.0, - "76": 1149693440.0, - "77": 1149693440.0, - "78": 1149693440.0, - "79": 1149693440.0, - "80": 1149693440.0, - "81": 1149693440.0, - "82": 1149693440.0, - "83": 1149693440.0, - "84": 1149693440.0, - "85": 1149693440.0, - "86": 1149693440.0, - "87": 1149693440.0, - "88": 1149693440.0, - "89": 1149693440.0, - "90": 1149693440.0, - "91": 1149693440.0, - "92": 1149693440.0, - "93": 1149693440.0, - "94": 1149693440.0, - "95": 1149693440.0, - "96": 1149693440.0, - "97": 1149693440.0, - "98": 1149693440.0, - "99": 1149693440.0, - "100": 1149693440.0 + "1": 974433792.0, + "2": 1143422976.0, + "3": 1143894016.0, + "4": 1147345408.0, + "5": 1147816448.0, + "6": 1147816448.0, + "7": 1148524544.0, + "8": 1148524544.0, + "9": 1148524544.0, + "10": 1148524544.0, + "11": 1148524544.0, + "12": 1148524544.0, + "13": 1148524544.0, + "14": 1148524544.0, + "15": 1148524544.0, + "16": 1148524544.0, + "17": 1148524544.0, + "18": 1148524544.0, + "19": 1148524544.0, + "20": 1148524544.0, + "21": 1148524544.0, + "22": 1148524544.0, + "23": 1148524544.0, + "24": 1148524544.0, + "25": 1148524544.0, + "26": 1149626368.0, + "27": 1149626368.0, + "28": 1149626368.0, + "29": 1149626368.0, + "30": 1149626368.0, + "31": 1149626368.0, + "32": 1149626368.0, + "33": 1149626368.0, + "34": 1149626368.0, + "35": 1149626368.0, + "36": 1149626368.0, + "37": 1149626368.0, + "38": 1149626368.0, + "39": 1149626368.0, + "40": 1149626368.0, + "41": 1149626368.0, + "42": 1149626368.0, + "43": 1149626368.0, + "44": 1149626368.0, + "45": 1149626368.0, + "46": 1149626368.0, + "47": 1149626368.0, + "48": 1149626368.0, + "49": 1149626368.0, + "50": 1149626368.0, + "51": 1149626368.0, + "52": 1149626368.0, + "53": 1149626368.0, + "54": 1149626368.0, + "55": 1149626368.0, + "56": 1149626368.0, + "57": 1149626368.0, + "58": 1149626368.0, + "59": 1149626368.0, + "60": 1149626368.0, + "61": 1149626368.0, + "62": 1149626368.0, + "63": 1149626368.0, + "64": 1149626368.0, + "65": 1149626368.0, + "66": 1149626368.0, + "67": 1149626368.0, + "68": 1149626368.0, + "69": 1149626368.0, + "70": 1149626368.0, + "71": 1149626368.0, + "72": 1149626368.0, + "73": 1149626368.0, + "74": 1149626368.0, + "75": 1149626368.0, + "76": 1149626368.0, + "77": 1149626368.0, + "78": 1149626368.0, + "79": 1149626368.0, + "80": 1149626368.0, + "81": 1149626368.0, + "82": 1149626368.0, + "83": 1149626368.0, + "84": 1149626368.0, + "85": 1149626368.0, + "86": 1149626368.0, + "87": 1149626368.0, + "88": 1149626368.0, + "89": 1149626368.0, + "90": 1149626368.0, + "91": 1149626368.0, + "92": 1149626368.0, + "93": 1149626368.0, + "94": 1149626368.0, + "95": 1149626368.0, + "96": 1149626368.0, + "97": 1149626368.0, + "98": 1149626368.0, + "99": 1149626368.0, + "100": 1149626368.0 } }, "iteration-time": { - "start_step": 1, + "start_step": 2, "end_step": 100, "step_interval": 1, "values": { - "1": "nan", - "2": 8.40786, - "3": 0.89447, - "4": 0.87487, - "5": 0.85695, - "6": 0.85548, - "7": 0.86384, - "8": 0.8398, - "9": 0.83442, - "10": 0.83457, - "11": 0.83165, - "12": 0.82049, - "13": 0.81938, - "14": 0.8372, - "15": 0.81635, - "16": 0.82269, - "17": 0.81755, - "18": 0.82139, - "19": 0.81834, - "20": 0.81571, - "21": 0.82027, - "22": 0.81783, - "23": 0.82434, - "24": 0.8179, - "25": 0.81779, - "26": 0.80609, - "27": 0.81441, - "28": 0.83081, - "29": 0.82504, - "30": 0.81873, - "31": 0.82454, - "32": 0.81663, - "33": 0.80909, - "34": 0.82198, - "35": 0.81846, - "36": 0.81614, - "37": 0.81026, - "38": 0.84604, - "39": 0.82085, - "40": 0.8318, - "41": 0.82267, - "42": 0.81837, - "43": 0.87684, - "44": 0.81896, - "45": 0.82655, - "46": 0.8241, - "47": 0.82308, - "48": 0.81433, - "49": 0.83989, - "50": 0.82395, - "51": 0.87417, - "52": 0.8737, - "53": 0.81483, - "54": 0.82825, - "55": 0.83667, - "56": 0.83546, - "57": 0.83562, - "58": 0.83505, - "59": 0.83375, - "60": 0.83021, - "61": 0.82875, - "62": 0.83214, - "63": 0.83746, - "64": 0.83687, - "65": 0.8281, - "66": 0.8317, - "67": 0.82752, - "68": 0.82693, - "69": 0.83293, - "70": 0.83375, - "71": 0.8272, - "72": 0.82716, - "73": 0.83134, - "74": 1.39559, - "75": 1.46874, - "76": 0.83059, - "77": 0.83236, - "78": 0.83428, - "79": 0.835, - "80": 0.83444, - "81": 0.83542, - "82": 0.84117, - "83": 0.83432, - "84": 0.82381, - "85": 0.831, - "86": 1.48456, - "87": 1.37924, - "88": 0.82167, - "89": 0.82408, - "90": 0.81692, - "91": 0.81059, - "92": 0.81301, - "93": 0.8096, - "94": 0.8091, - "95": 0.80549, - "96": 0.80731, - "97": 0.81231, - "98": 0.8007, - "99": 0.80887, - "100": 0.81095 + "2": 8.99025, + "3": 0.90833, + "4": 0.86431, + "5": 0.87482, + "6": 0.8573, + "7": 0.85553, + "8": 0.83513, + "9": 0.83483, + "10": 0.83557, + "11": 0.83528, + "12": 1.58117, + "13": 0.82194, + "14": 0.81823, + "15": 0.81808, + "16": 0.81282, + "17": 0.82548, + "18": 0.81502, + "19": 0.81167, + "20": 0.81094, + "21": 0.83617, + "22": 0.828, + "23": 0.82514, + "24": 0.85341, + "25": 0.81784, + "26": 0.81255, + "27": 0.81988, + "28": 0.84249, + "29": 1.54481, + "30": 0.82635, + "31": 0.81779, + "32": 0.83092, + "33": 0.82788, + "34": 0.8237, + "35": 0.83024, + "36": 0.81681, + "37": 0.81326, + "38": 1.61795, + "39": 0.84407, + "40": 0.85127, + "41": 0.82922, + "42": 0.83611, + "43": 0.81901, + "44": 2.40554, + "45": 0.81924, + "46": 0.84478, + "47": 0.8247, + "48": 1.5457, + "49": 0.81497, + "50": 0.80868, + "51": 0.95984, + "52": 0.95192, + "53": 0.8168, + "54": 0.83474, + "55": 0.82917, + "56": 0.81693, + "57": 0.81814, + "58": 0.80755, + "59": 0.80597, + "60": 0.81208, + "61": 0.81909, + "62": 0.80757, + "63": 0.8268, + "64": 0.81229, + "65": 0.80745, + "66": 0.81789, + "67": 0.79693, + "68": 0.81825, + "69": 0.82426, + "70": 0.8204, + "71": 0.80984, + "72": 0.8026, + "73": 0.80831, + "74": 0.86362, + "75": 0.82144, + "76": 0.87744, + "77": 0.81254, + "78": 0.82353, + "79": 0.85476, + "80": 0.8258, + "81": 0.81312, + "82": 0.81209, + "83": 0.80825, + "84": 0.81139, + "85": 0.80849, + "86": 0.81299, + "87": 0.81396, + "88": 0.80619, + "89": 0.79982, + "90": 0.80843, + "91": 0.81705, + "92": 0.8122, + "93": 0.8039, + "94": 0.80977, + "95": 0.81732, + "96": 0.80769, + "97": 0.81238, + "98": 0.80923, + "99": 0.80613, + "100": 0.80743 } } -} \ No newline at end of file +} diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/model_config.yaml index efdac2478fd..58540a9d6e7 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/model_config.yaml @@ -96,3 +96,4 @@ METRICS: # - "mem-allocated-bytes" # - "mem-max-allocated-bytes" - "mtp_1 loss" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph_1node/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph_1node/model_config.yaml index 26af6497637..90011c10c76 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph_1node/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph_1node/model_config.yaml @@ -95,3 +95,4 @@ METRICS: # - "mem-allocated-bytes" # - "mem-max-allocated-bytes" - "mtp_1 loss" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_cuda_graphs_pad_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_cuda_graphs_pad_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml index 3bd326a56e1..2355f67a22f 100644 --- a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_cuda_graphs_pad_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_cuda_graphs_pad_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml @@ -1,5 +1,6 @@ ENV_VARS: CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE: 1 NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 NCCL_ALGO: Ring CUBLAS_WORKSPACE_CONFIG: :4096:8 @@ -86,3 +87,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp2_pp2_ep2_gptoss_20b_swa/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp2_pp2_ep2_gptoss_20b_swa/model_config.yaml index 7d87f0a9998..facc8c97da3 100644 --- a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp2_pp2_ep2_gptoss_20b_swa/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp2_pp2_ep2_gptoss_20b_swa/model_config.yaml @@ -102,3 +102,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_cudagraph_zmq/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_cudagraph_zmq/model_config.yaml index 80e2a37c250..a4131cb245b 100644 --- a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_cudagraph_zmq/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_cudagraph_zmq/model_config.yaml @@ -1,5 +1,6 @@ ENV_VARS: CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE: 1 NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 NCCL_ALGO: Ring CUBLAS_WORKSPACE_CONFIG: :4096:8 @@ -87,3 +88,4 @@ METRICS: - "generated_tokens" - "logprobs" - "routing_indices" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_zmq/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_zmq/model_config.yaml index 479cb7a4751..b32eb7e76dc 100644 --- a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_zmq/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_zmq/model_config.yaml @@ -84,3 +84,4 @@ METRICS: - "generated_tokens" - "logprobs" - "routing_indices" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_zmq_suspend_resume/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_zmq_suspend_resume/model_config.yaml index 1f302455440..5e63740403a 100644 --- a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_zmq_suspend_resume/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_etp1_pp1_ep8_16B_logitsmatch_zmq_suspend_resume/model_config.yaml @@ -87,3 +87,4 @@ MODEL_ARGS: --no-rl-persist-cuda-graphs: true METRICS: +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_chunked_prefill/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_chunked_prefill/model_config.yaml index db20ea13cf1..35f0b7951e7 100644 --- a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_chunked_prefill/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_chunked_prefill/model_config.yaml @@ -83,3 +83,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml index 5ed1f1205f6..ea8eb4c2f03 100644 --- a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml @@ -80,3 +80,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_prefix_caching/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_prefix_caching/model_config.yaml index d293646fa2b..69fba5c46cc 100644 --- a/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_prefix_caching/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_dynamic_inference_tp4_pp1_ep4_16B_prefix_caching/model_config.yaml @@ -81,3 +81,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_grpo_tp8tp4_pp1_ep8ep2_dp8_throughputtest/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_grpo_tp8tp4_pp1_ep8ep2_dp8_throughputtest/model_config.yaml index 22cc8d5e4d2..2e9157cbcc2 100644 --- a/tests/functional_tests/test_cases/moe/gpt_grpo_tp8tp4_pp1_ep8ep2_dp8_throughputtest/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_grpo_tp8tp4_pp1_ep8ep2_dp8_throughputtest/model_config.yaml @@ -136,3 +136,4 @@ METRICS: - "num-zeros" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_static_inference_cuda_graphs_pad_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_static_inference_cuda_graphs_pad_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml index 049c9090099..013eee2e49f 100644 --- a/tests/functional_tests/test_cases/moe/gpt_static_inference_cuda_graphs_pad_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_static_inference_cuda_graphs_pad_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml @@ -85,3 +85,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_static_inference_tp1_pp1_ep1_16B_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_static_inference_tp1_pp1_ep1_16B_logitsmatch/model_config.yaml index 03895d97ee9..9be551d93ea 100644 --- a/tests/functional_tests/test_cases/moe/gpt_static_inference_tp1_pp1_ep1_16B_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_static_inference_tp1_pp1_ep1_16B_logitsmatch/model_config.yaml @@ -1,5 +1,6 @@ ENV_VARS: CUDA_DEVICE_MAX_CONNECTIONS: 1 + NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE: 1 NVTE_ALLOW_NONDETERMINISTIC_ALGO: 0 NCCL_ALGO: Ring CUBLAS_WORKSPACE_CONFIG: :4096:8 @@ -79,3 +80,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/moe/gpt_static_inference_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml b/tests/functional_tests/test_cases/moe/gpt_static_inference_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml index 9259d63c9d1..93532c83707 100644 --- a/tests/functional_tests/test_cases/moe/gpt_static_inference_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml +++ b/tests/functional_tests/test_cases/moe/gpt_static_inference_tp4_pp1_ep4_16B_logitsmatch/model_config.yaml @@ -80,3 +80,4 @@ MODEL_ARGS: METRICS: - "generated_tokens" - "logprobs" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/multimodal-llava/multimodal_llava_mcore_te_tp1_pp1/model_config.yaml b/tests/functional_tests/test_cases/multimodal-llava/multimodal_llava_mcore_te_tp1_pp1/model_config.yaml index 377371b2370..2fdd47c8def 100644 --- a/tests/functional_tests/test_cases/multimodal-llava/multimodal_llava_mcore_te_tp1_pp1/model_config.yaml +++ b/tests/functional_tests/test_cases/multimodal-llava/multimodal_llava_mcore_te_tp1_pp1/model_config.yaml @@ -51,3 +51,4 @@ MODEL_ARGS: --mock-data: true --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/multimodal-llava/multimodal_llava_mcore_te_tp4_sp_cp2/model_config.yaml b/tests/functional_tests/test_cases/multimodal-llava/multimodal_llava_mcore_te_tp4_sp_cp2/model_config.yaml index 9332c934613..ac4f20b2781 100644 --- a/tests/functional_tests/test_cases/multimodal-llava/multimodal_llava_mcore_te_tp4_sp_cp2/model_config.yaml +++ b/tests/functional_tests/test_cases/multimodal-llava/multimodal_llava_mcore_te_tp4_sp_cp2/model_config.yaml @@ -57,3 +57,4 @@ MODEL_ARGS: --log-memory-to-tensorboard: true --calculate-per-token-loss: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_gb200/model_config.yaml b/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_gb200/model_config.yaml index 63d6eb050fc..c7b9da95622 100644 --- a/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_gb200/model_config.yaml +++ b/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_gb200/model_config.yaml @@ -160,3 +160,4 @@ METRICS: - "lm loss" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_gb200_sm/model_config.yaml b/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_gb200_sm/model_config.yaml index 87d9f587ee1..75dbca82700 100644 --- a/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_gb200_sm/model_config.yaml +++ b/tests/functional_tests/test_cases/nemotron/nemotron3_super_release_gb200_sm/model_config.yaml @@ -160,3 +160,4 @@ METRICS: - "lm loss" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_11b_mcore_tp4_pp1/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_11b_mcore_tp4_pp1/model_config.yaml index 8dc5cf5a713..687ee924f5b 100644 --- a/tests/functional_tests/test_cases/t5/t5_11b_mcore_tp4_pp1/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_11b_mcore_tp4_pp1/model_config.yaml @@ -57,3 +57,4 @@ METRICS: - "num-zeros" - "mem-allocated-bytes" - "mem-max-allocated-bytes" +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_mcore_te_tp1_pp1_vp1_resume_torch/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_mcore_te_tp1_pp1_vp1_resume_torch/model_config.yaml index 00a11b9a439..6adfb92fb77 100644 --- a/tests/functional_tests/test_cases/t5/t5_mcore_te_tp1_pp1_vp1_resume_torch/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_mcore_te_tp1_pp1_vp1_resume_torch/model_config.yaml @@ -60,3 +60,4 @@ MODEL_ARGS: # the worker-queue hang surface; test data is tiny so perf impact is nil. --num-workers: 0 TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_mcore_te_tp2_pp1_vp1/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_mcore_te_tp2_pp1_vp1/model_config.yaml index 277d637cb4d..c872de4b25b 100644 --- a/tests/functional_tests/test_cases/t5/t5_mcore_te_tp2_pp1_vp1/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_mcore_te_tp2_pp1_vp1/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_mcore_te_tp2_pp1_vp1_sequence_parallel/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_mcore_te_tp2_pp1_vp1_sequence_parallel/model_config.yaml index 8463c0ebfa4..b58764cf414 100644 --- a/tests/functional_tests/test_cases/t5/t5_mcore_te_tp2_pp1_vp1_sequence_parallel/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_mcore_te_tp2_pp1_vp1_sequence_parallel/model_config.yaml @@ -55,3 +55,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_mcore_te_tp4_pp1/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_mcore_te_tp4_pp1/model_config.yaml index f4a4e7c3b1c..bc2837c7fa3 100644 --- a/tests/functional_tests/test_cases/t5/t5_mcore_te_tp4_pp1/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_mcore_te_tp4_pp1/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_mcore_te_tp4_pp1_resume_torch_dist/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_mcore_te_tp4_pp1_resume_torch_dist/model_config.yaml index 5f170ec8f43..dea7ae55995 100644 --- a/tests/functional_tests/test_cases/t5/t5_mcore_te_tp4_pp1_resume_torch_dist/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_mcore_te_tp4_pp1_resume_torch_dist/model_config.yaml @@ -54,3 +54,4 @@ MODEL_ARGS: --attention-backend: unfused --log-memory-to-tensorboard: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_mcore_tp1_pp1_vp1/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_mcore_tp1_pp1_vp1/model_config.yaml index 805cc3e9858..b5e2398a84e 100644 --- a/tests/functional_tests/test_cases/t5/t5_mcore_tp1_pp1_vp1/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_mcore_tp1_pp1_vp1/model_config.yaml @@ -58,3 +58,4 @@ MODEL_ARGS: # the worker-queue hang surface; test data is tiny so perf impact is nil. --num-workers: 0 TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_mcore_tp1_pp1_vp1_resume_torch/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_mcore_tp1_pp1_vp1_resume_torch/model_config.yaml index 5b1e192a3a9..5e4f1aa8019 100644 --- a/tests/functional_tests/test_cases/t5/t5_mcore_tp1_pp1_vp1_resume_torch/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_mcore_tp1_pp1_vp1_resume_torch/model_config.yaml @@ -58,3 +58,4 @@ MODEL_ARGS: # the worker-queue hang surface; test data is tiny so perf impact is nil. --num-workers: 0 TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_mcore_tp2_pp1_vp1/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_mcore_tp2_pp1_vp1/model_config.yaml index df1bb9d1833..b7391532152 100644 --- a/tests/functional_tests/test_cases/t5/t5_mcore_tp2_pp1_vp1/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_mcore_tp2_pp1_vp1/model_config.yaml @@ -52,3 +52,4 @@ MODEL_ARGS: --ckpt-format: torch --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_mcore_tp4_pp1/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_mcore_tp4_pp1/model_config.yaml index 33d79798194..a0d1d0a40d8 100644 --- a/tests/functional_tests/test_cases/t5/t5_mcore_tp4_pp1/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_mcore_tp4_pp1/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --dist-ckpt-strictness: log_all # backward compatibility for TE changes --log-memory-to-tensorboard: true TEST_TYPE: regular +LAUNCHER: ft_launcher diff --git a/tests/functional_tests/test_cases/t5/t5_mcore_tp4_pp1_resume_torch_dist/model_config.yaml b/tests/functional_tests/test_cases/t5/t5_mcore_tp4_pp1_resume_torch_dist/model_config.yaml index 15096f11362..44aa3101027 100644 --- a/tests/functional_tests/test_cases/t5/t5_mcore_tp4_pp1_resume_torch_dist/model_config.yaml +++ b/tests/functional_tests/test_cases/t5/t5_mcore_tp4_pp1_resume_torch_dist/model_config.yaml @@ -53,3 +53,4 @@ MODEL_ARGS: --dist-ckpt-strictness: log_all # backward compatibility for TE changes --log-memory-to-tensorboard: true TEST_TYPE: ckpt-resume +LAUNCHER: ft_launcher diff --git a/tests/performance_tests/client/static_benchmark.py b/tests/performance_tests/client/static_benchmark.py index 10c945c8afc..1c2d64e9110 100644 --- a/tests/performance_tests/client/static_benchmark.py +++ b/tests/performance_tests/client/static_benchmark.py @@ -2,8 +2,9 @@ """Static throughput/latency benchmark against an OpenAI-compatible completions server. Fires --batch-size requests simultaneously via asyncio.gather, waits for all to -finish, and reports throughput, latency (avg/p50/p99), and TPOT. Iterates over -warmup + timed batches and emits a JSON results file consumable by +finish, and reports throughput, latency (avg/p50/p99), and TPOT. Warmup batches +are widened to cover every data-parallel worker when needed. Timed batches keep +the requested batch size and emit a JSON results file consumable by compare_to_baseline.py. Hits the server's POST /v1/completions endpoint directly via aiohttp — no @@ -92,10 +93,15 @@ async def _run_batch( url: str, prompts: list[str], iter_start_index: int, + request_count: int | None = None, ) -> tuple[list[int], list[int], list[float], float]: - """Fire batch_size requests in parallel. Cycles through `prompts` deterministically - starting at `iter_start_index` so each timed iteration sees the same prompt - distribution (reduces run-to-run variance for gsm8k mode).""" + """Fire requests in parallel, defaulting to the measured batch size. + + Cycles through `prompts` deterministically starting at `iter_start_index` + so each timed iteration sees the same prompt distribution (reduces + run-to-run variance for gsm8k mode). + """ + request_count = args.batch_size if request_count is None else request_count t0 = time.perf_counter() results = await asyncio.gather( *[ @@ -107,7 +113,7 @@ async def _run_batch( args.num_output_tokens, args.temperature, ) - for i in range(args.batch_size) + for i in range(request_count) ] ) wall = time.perf_counter() - t0 @@ -122,8 +128,14 @@ def _percentile(sorted_values: list[float], pct: float) -> float: return sorted_values[idx] +def _get_warmup_batch_size(batch_size: int, data_parallel_size: int) -> int: + """Keep batch-shape warmup while issuing enough requests to cover DP workers.""" + return max(batch_size, data_parallel_size) + + async def main(args: argparse.Namespace) -> dict: url = f"{args.server_url.rstrip('/')}/completions" + warmup_batch_size = _get_warmup_batch_size(args.batch_size, args.data_parallel_size) if args.dataset == "gsm8k": prompts = _load_gsm8k_prompts() @@ -140,26 +152,34 @@ async def main(args: argparse.Namespace) -> dict: print(f"Dataset : {prompt_source}") print(f"Output tokens : {args.num_output_tokens}") print(f"Warmup iters : {args.num_warmup_iters}") + print(f"Warmup batch : {warmup_batch_size}") print(f"Timed iters : {args.num_iters}", flush=True) connector = aiohttp.TCPConnector(limit=0) async with aiohttp.ClientSession(connector=connector) as session: - cursor = 0 + warmup_cursor = 0 for i in range(args.num_warmup_iters): - print(f"\nWarmup {i + 1}/{args.num_warmup_iters}...", flush=True) - await _run_batch(session, args, url, prompts, cursor) - cursor += args.batch_size + print( + f"\nWarmup {i + 1}/{args.num_warmup_iters} (batch={warmup_batch_size})...", + flush=True, + ) + await _run_batch( + session, args, url, prompts, warmup_cursor, request_count=warmup_batch_size + ) + warmup_cursor += warmup_batch_size all_wall: list[float] = [] all_output_tokens: list[int] = [] all_input_tokens: list[int] = [] all_latencies: list[float] = [] + # Keep the timed prompt sequence stable when widening warmup batches. + timed_cursor = args.num_warmup_iters * args.batch_size for i in range(args.num_iters): input_counts, output_counts, latencies, wall = await _run_batch( - session, args, url, prompts, cursor + session, args, url, prompts, timed_cursor ) - cursor += args.batch_size + timed_cursor += args.batch_size total_out = sum(output_counts) all_wall.append(wall) all_output_tokens.append(total_out) @@ -226,6 +246,13 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--num-output-tokens", type=int, default=128) parser.add_argument("--temperature", type=float, default=0.0) parser.add_argument("--num-warmup-iters", type=int, default=2) + parser.add_argument( + "--data-parallel-size", + type=int, + default=1, + help="Number of coordinator-addressable data-parallel workers. Warmup batches " + "use at least this many concurrent requests; timed batch size is unchanged.", + ) parser.add_argument("--num-iters", type=int, default=5) parser.add_argument( "--output-json", diff --git a/tests/performance_tests/shell_test_utils/determinism/print_nsys_leaderboard.py b/tests/performance_tests/shell_test_utils/determinism/print_nsys_leaderboard.py index da6a050d2d4..a95d976720d 100644 --- a/tests/performance_tests/shell_test_utils/determinism/print_nsys_leaderboard.py +++ b/tests/performance_tests/shell_test_utils/determinism/print_nsys_leaderboard.py @@ -11,7 +11,7 @@ import sys from pathlib import Path -MAX_DET_NONDET_RATIO = 1.25 +MAX_DET_NONDET_RATIO = 1.35 MEASUREMENT_ITER = 5 # steady-state; iter 7 is noisy under nsys profile teardown LEADERBOARD_TOP_N = 20 # Strip per-call-site ``, op_id = N`` and autograd-engine ``, seq = N`` so diff --git a/tests/performance_tests/shell_test_utils/run_perf_test.sh b/tests/performance_tests/shell_test_utils/run_perf_test.sh index b8effb030b5..830af93c5e0 100755 --- a/tests/performance_tests/shell_test_utils/run_perf_test.sh +++ b/tests/performance_tests/shell_test_utils/run_perf_test.sh @@ -88,6 +88,12 @@ NUM_TIMED_ITERS=$("$YQ" '.NUM_TIMED_ITERS // 5' "$CONFIG_PATH") # hybrid models should use 'gsm8k' — synthetic input gives misleading # perf because every token is identical (uniform expert routing, hot KV). DATASET=$("$YQ" '.DATASET // "synthetic"' "$CONFIG_PATH") +# Async prefill scheduling for dynamic batching. When true, the server is +# launched with --inference-dynamic-batching-async-sched-mode async (overlaps +# the prefill scheduler with GPU compute). Requires greedy sampling / no +# logprobs / no stop words — the static benchmark client already satisfies +# these (temperature 0.0, ignore_eos, no stop tokens, no logprobs requested). +ASYNC_SCHED=$("$YQ" '.ASYNC_SCHED // false' "$CONFIG_PATH") mapfile -t BATCH_SIZES < <("$YQ" '.BATCH_SIZES[]' "$CONFIG_PATH") # For MoE configs, expert-parallelism is orthogonal to DP and reshapes the @@ -95,6 +101,9 @@ mapfile -t BATCH_SIZES < <("$YQ" '.BATCH_SIZES[]' "$CONFIG_PATH") # and MoE-with-DP=1 picks up EP correctly. GROUP_SIZE=$((DP > EP ? DP : EP)) WORLD_SIZE=$((TP * PP * GROUP_SIZE)) +# The inference coordinator uses the dense-model DP group. EP-only configs +# therefore expose GROUP_SIZE workers even when the YAML DP value is one. +COORDINATOR_WORKERS=$GROUP_SIZE ARGS_FILE="$PERF_DIR/server/model_args/${MODEL}.args" if [[ ! -f "$ARGS_FILE" ]]; then echo "[run_perf_test] error: model args file $ARGS_FILE not found" >&2 @@ -102,6 +111,7 @@ if [[ ! -f "$ARGS_FILE" ]]; then fi echo "[run_perf_test] MODEL=$MODEL TP=$TP PP=$PP DP=$DP EP=$EP world_size=$WORLD_SIZE dataset=$DATASET" +echo "[run_perf_test] coordinator workers: $COORDINATOR_WORKERS" echo "[run_perf_test] ISL=$NUM_INPUT_TOKENS OSL=$NUM_OUTPUT_TOKENS" echo "[run_perf_test] batch sizes: ${BATCH_SIZES[*]}" @@ -190,6 +200,20 @@ SERVER_COMMON_ARGS=( --host 0.0.0.0 ) +# Enable async prefill scheduling when the test case opts in. Async scheduling +# requires materialize_only_last_token_logits=True. run_dynamic_text_generation_server +# force-sets return_log_probs=True (for echo/loglikelihood support), which would +# flip materialize_only_last_token_logits to False; passing --skip-prompt-log-probs +# keeps it True (materialize = not(return_log_probs and not skip_prompt_log_probs)). +# The perf client never requests prompt logprobs, so this is a no-op for the metrics. +if [[ "$ASYNC_SCHED" == "true" ]]; then + echo "[run_perf_test] async scheduling enabled (--inference-dynamic-batching-async-sched-mode async --skip-prompt-log-probs)" + SERVER_COMMON_ARGS+=( + --inference-dynamic-batching-async-sched-mode async + --skip-prompt-log-probs + ) +fi + ( cd "$ROOT_DIR" uv run --no-sync python -m torch.distributed.run \ @@ -257,6 +281,7 @@ for BS in "${BATCH_SIZES[@]}"; do --num-input-tokens "$NUM_INPUT_TOKENS" \ --num-output-tokens "$NUM_OUTPUT_TOKENS" \ --num-warmup-iters "$NUM_WARMUP_ITERS" \ + --data-parallel-size "$COORDINATOR_WORKERS" \ --num-iters "$NUM_TIMED_ITERS" \ --output-json "$RESULTS_JSON" \ 2>&1 | tee -a "$RESULTS_ROOT/benchmark.log" diff --git a/tests/performance_tests/test_cases/gpt/gpt_583m_perf/baseline_values.json b/tests/performance_tests/test_cases/gpt/gpt_583m_perf/baseline_values.json index 5c83450644f..be494f2fcf7 100644 --- a/tests/performance_tests/test_cases/gpt/gpt_583m_perf/baseline_values.json +++ b/tests/performance_tests/test_cases/gpt/gpt_583m_perf/baseline_values.json @@ -2,47 +2,51 @@ "h100": { "batch_1": { "batch_size": 1, - "num_input_tokens": 512, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, "num_output_tokens": 128, "num_iters": 5, - "throughput_tok_per_sec": 22.08203581766323, - "avg_latency_ms": 5796.511190757155, - "p50_latency_ms": 5842.958671972156, - "p99_latency_ms": 5919.582245871425, - "tpot_ms_per_tok": 45.28567964734975 + "throughput_tok_per_sec": 44.399582392645556, + "avg_latency_ms": 2882.8655768185854, + "p50_latency_ms": 2883.104130625725, + "p99_latency_ms": 2886.1329462379217, + "tpot_ms_per_tok": 22.522734361700714 }, "batch_8": { "batch_size": 8, - "num_input_tokens": 512, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, "num_output_tokens": 128, "num_iters": 5, - "throughput_tok_per_sec": 357.49352020243185, - "avg_latency_ms": 2808.0778209026903, - "p50_latency_ms": 2813.233459368348, - "p99_latency_ms": 2896.668652072549, - "tpot_ms_per_tok": 22.378027986269444 + "throughput_tok_per_sec": 343.99702071382734, + "avg_latency_ms": 2918.7685920856893, + "p50_latency_ms": 2911.766432225704, + "p99_latency_ms": 3012.841146439314, + "tpot_ms_per_tok": 23.25601536722388 }, "batch_32": { "batch_size": 32, - "num_input_tokens": 512, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, "num_output_tokens": 128, "num_iters": 5, - "throughput_tok_per_sec": 1432.3391490750541, - "avg_latency_ms": 2812.452972715255, - "p50_latency_ms": 2819.5638693869114, - "p99_latency_ms": 2865.5092362314463, - "tpot_ms_per_tok": 22.341077544842847 + "throughput_tok_per_sec": 1378.968944067799, + "avg_latency_ms": 2920.1371596893296, + "p50_latency_ms": 2913.53677585721, + "p99_latency_ms": 2975.6032899022102, + "tpot_ms_per_tok": 23.2057437824551 }, "batch_128": { "batch_size": 128, - "num_input_tokens": 512, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, "num_output_tokens": 128, "num_iters": 5, - "throughput_tok_per_sec": 5643.249306634135, - "avg_latency_ms": 2839.432980850688, - "p50_latency_ms": 2846.7628210783005, - "p99_latency_ms": 2900.5435090512037, - "tpot_ms_per_tok": 22.681967966491356 + "throughput_tok_per_sec": 5414.296145922958, + "avg_latency_ms": 2964.334140153369, + "p50_latency_ms": 2969.226948916912, + "p99_latency_ms": 3022.6969085633755, + "tpot_ms_per_tok": 23.641115400823765 } } } diff --git a/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched/baseline_values.json b/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched/baseline_values.json new file mode 100644 index 00000000000..026a7354660 --- /dev/null +++ b/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched/baseline_values.json @@ -0,0 +1,52 @@ +{ + "h100": { + "batch_1": { + "batch_size": 1, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, + "num_output_tokens": 128, + "num_iters": 5, + "throughput_tok_per_sec": 42.75601962213067, + "avg_latency_ms": 2993.6806965619326, + "p50_latency_ms": 2992.9271759465337, + "p99_latency_ms": 3003.761636093259, + "tpot_ms_per_tok": 23.388519530999474 + }, + "batch_8": { + "batch_size": 8, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, + "num_output_tokens": 128, + "num_iters": 5, + "throughput_tok_per_sec": 344.0237043263471, + "avg_latency_ms": 2920.4657254274935, + "p50_latency_ms": 2921.021580696106, + "p99_latency_ms": 2982.457813806832, + "tpot_ms_per_tok": 23.254211554012727 + }, + "batch_32": { + "batch_size": 32, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, + "num_output_tokens": 128, + "num_iters": 5, + "throughput_tok_per_sec": 1372.033198644746, + "avg_latency_ms": 2927.9737344244495, + "p50_latency_ms": 2925.4276445135474, + "p99_latency_ms": 3030.2867460995913, + "tpot_ms_per_tok": 23.32305080635706 + }, + "batch_128": { + "batch_size": 128, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, + "num_output_tokens": 128, + "num_iters": 5, + "throughput_tok_per_sec": 5395.836458577838, + "avg_latency_ms": 2958.10645565507, + "p50_latency_ms": 2953.0187863856554, + "p99_latency_ms": 3037.2949857264757, + "tpot_ms_per_tok": 23.72199398232624 + } + } +} diff --git a/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched/model_config.yaml b/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched/model_config.yaml new file mode 100644 index 00000000000..d95f385ebe4 --- /dev/null +++ b/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched/model_config.yaml @@ -0,0 +1,38 @@ +# Inference perf test: 583M mcore-mistral checkpoint, TP=1 PP=1 DP=8 (8 GPUs), +# async prefill scheduling enabled. +# +# Same 583M checkpoint / DP=8 dynamic-batching path as gpt_583m_perf, but the +# server runs with --inference-dynamic-batching-async-sched-mode async (the +# scheduler overlaps prefill setup with GPU compute). Measures the throughput / +# latency of the async scheduling path exercised by the functional test +# gpt_dynamic_inference_tp1_pp1_583m_async_sched. +# +# Async scheduling requires greedy sampling / no logprobs / no stop words; the +# static benchmark client already satisfies these (temperature 0.0, ignore_eos). +# +# Baseline values are recorded by `RECORD_BASELINE=1 run_perf_test.sh ...` +# and compared on subsequent runs with TOLERANCE_PCT tolerance. + +MODEL: gpt_583m +TP: 1 +PP: 1 +DP: 8 +ASYNC_SCHED: true +NUM_INPUT_TOKENS: 512 +NUM_OUTPUT_TOKENS: 128 +NUM_WARMUP_ITERS: 2 +NUM_TIMED_ITERS: 5 +BATCH_SIZES: + - 1 + - 8 + - 32 + - 128 +TOLERANCE_PCT: 10 +# p99 omitted on purpose: with NUM_TIMED_ITERS=5 it is the max of 5 samples, +# not a real percentile, so it produces flaky regressions even when throughput +# / avg / p50 are stable. p99 is still recorded in results.json for visibility. +METRICS: + - throughput_tok_per_sec + - avg_latency_ms + - p50_latency_ms + - tpot_ms_per_tok diff --git a/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched_gb200_4gpu/baseline_values.json b/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched_gb200_4gpu/baseline_values.json new file mode 100644 index 00000000000..122e3399183 --- /dev/null +++ b/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched_gb200_4gpu/baseline_values.json @@ -0,0 +1,40 @@ +{ + "gb200": { + "batch_8": { + "batch_size": 8, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, + "num_output_tokens": 128, + "num_iters": 5, + "throughput_tok_per_sec": 306.77571598122546, + "avg_latency_ms": 3292.265848722309, + "p50_latency_ms": 3284.7761889570393, + "p99_latency_ms": 3369.791687990073, + "tpot_ms_per_tok": 26.07768341249539 + }, + "batch_32": { + "batch_size": 32, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, + "num_output_tokens": 128, + "num_iters": 5, + "throughput_tok_per_sec": 1217.121447486444, + "avg_latency_ms": 3322.6687765843963, + "p50_latency_ms": 3319.4968919851817, + "p99_latency_ms": 3425.9593110182323, + "tpot_ms_per_tok": 26.29154228288826 + }, + "batch_128": { + "batch_size": 128, + "dataset": "synthetic", + "num_input_tokens_avg": 512.0, + "num_output_tokens": 128, + "num_iters": 5, + "throughput_tok_per_sec": 4573.446310856698, + "avg_latency_ms": 3390.523879592547, + "p50_latency_ms": 3371.0599309997633, + "p99_latency_ms": 3771.179543051403, + "tpot_ms_per_tok": 27.987646798464993 + } + } +} diff --git a/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched_gb200_4gpu/model_config.yaml b/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched_gb200_4gpu/model_config.yaml new file mode 100644 index 00000000000..16bc84ab5c5 --- /dev/null +++ b/tests/performance_tests/test_cases/gpt/gpt_583m_perf_async_sched_gb200_4gpu/model_config.yaml @@ -0,0 +1,36 @@ +# Inference perf test: 583M mcore-mistral checkpoint, TP=1 PP=1 DP=4 (4 GPUs), +# async prefill scheduling enabled. +# +# GB200 single-node variant of gpt_583m_perf_async_sched. GB200 nodes have a +# 4-GPU/node limit, so the DP=8 configuration used on H100 doesn't fit on a +# single GB200 node. This is a separate test (different world size, different +# baseline) — not a multi-node port of the DP=8 case. +# +# The server runs with --inference-dynamic-batching-async-sched-mode async. +# Async scheduling requires greedy sampling / no logprobs / no stop words; the +# static benchmark client already satisfies these (temperature 0.0, ignore_eos). + +MODEL: gpt_583m +TP: 1 +PP: 1 +DP: 4 +ASYNC_SCHED: true +NUM_INPUT_TOKENS: 512 +NUM_OUTPUT_TOKENS: 128 +# 2 warmup iters can leave the first timed iteration cold on GB200, poisoning +# the mean/tail metrics; warm up more so timing starts at steady state. +NUM_WARMUP_ITERS: 5 +NUM_TIMED_ITERS: 5 +# batch_size=1 is omitted: at single-stream the 4-GPU (DP=4) deployment is +# barely utilized, so per-iteration jitter dominates and the small-sample p50 +# latency is too noisy to gate on. Larger batches are stable perf signals. +BATCH_SIZES: + - 8 + - 32 + - 128 +TOLERANCE_PCT: 10 +METRICS: + - throughput_tok_per_sec + - avg_latency_ms + - p50_latency_ms + - tpot_ms_per_tok diff --git a/tests/test_utils/python_scripts/generate_jet_trigger_job.py b/tests/test_utils/python_scripts/generate_jet_trigger_job.py index aca23ad9fee..e3b471a629d 100644 --- a/tests/test_utils/python_scripts/generate_jet_trigger_job.py +++ b/tests/test_utils/python_scripts/generate_jet_trigger_job.py @@ -9,6 +9,26 @@ from tests.test_utils.python_scripts import recipe_parser BASE_PATH = pathlib.Path(__file__).parent.resolve() +TRIAGE_LOG_PATH = "jet_workload.log" +TRIAGE_REPORT_PATH = "error_report.json" + + +def build_test_script(command: str) -> str: + """Wrap a workload command with non-blocking error extraction.""" + return "\n".join( + [ + "set +e", + "set -o pipefail", + f"{command} 2>&1 | tee {TRIAGE_LOG_PATH}", + 'exit_code=${PIPESTATUS[0]}', + "set -e", + ( + f"extract-errors {TRIAGE_LOG_PATH} --output {TRIAGE_REPORT_PATH} " + '--exit-code "$exit_code" || true' + ), + 'exit "$exit_code"', + ] + ) @click.command() @@ -70,6 +90,11 @@ "Empty/unset disables the cadence filter." ), ) +@click.option( + "--enable-error-extraction/--no-enable-error-extraction", + default=False, + help="Extract a structured error report from GitLab child-job output.", +) def main( scope: str, environment: str, @@ -91,7 +116,8 @@ def main( enable_lightweight_mode: bool = False, enable_warmup: Optional[bool] = None, cadence: Optional[str] = None, -): + enable_error_extraction: bool = False, +) -> None: # Treat empty string as "no cadence filter" so callers can wire shell # variables in directly without conditional flag emission. cadence_arg = cadence or None @@ -217,14 +243,20 @@ def main( elif warmup_job != "": needs.append({"job": warmup_job}) + test_script = " ".join(script) + artifact_paths = ["results/"] + if enable_error_extraction: + test_script = build_test_script(test_script) + artifact_paths.extend([TRIAGE_LOG_PATH, TRIAGE_REPORT_PATH]) + gitlab_pipeline[test_case['spec']['test_case']] = { "stage": f"{test_case['spec']['model']}", "image": f"{container_image}:{container_tag}", "tags": job_tags, "timeout": "7 days", "needs": needs, - "script": [" ".join(script)], - "artifacts": {"paths": ["results/"], "when": "always"}, + "script": [test_script], + "artifacts": {"paths": artifact_paths, "when": "always"}, "allow_failure": test_case["spec"].get("allow_failure", False) or test_case["spec"]["model"] == "gpt-nemo", "retry": { diff --git a/tests/test_utils/python_scripts/launch_jet_workload.py b/tests/test_utils/python_scripts/launch_jet_workload.py index ff79f88bc1f..f09eac7ba10 100644 --- a/tests/test_utils/python_scripts/launch_jet_workload.py +++ b/tests/test_utils/python_scripts/launch_jet_workload.py @@ -6,7 +6,6 @@ import pathlib import re import signal -import subprocess import sys import time import uuid @@ -33,44 +32,6 @@ logger = logging.getLogger(__name__) -def send_slack_alert(test_case: str, context: str, n_iteration: int, n_attempts: int) -> None: - """Send a Slack alert via notify.py for the current release pipeline state. - - Args: - test_case: Name of the release test case being run. - context: Human-readable context string appended to the pipeline context label. - n_iteration: Current training iteration (pipeline relaunch count). - n_attempts: Current attempt count within this iteration. - """ - pipeline_id = os.getenv("PARENT_PIPELINE_ID") - pipeline_created_at = os.getenv("CI_PIPELINE_CREATED_AT", "") - - if not pipeline_id or not pipeline_created_at: - logger.info("Missing PARENT_PIPELINE_ID or CI_PIPELINE_CREATED_AT, skipping Slack alert.") - return - - pipeline_context = f"{test_case} | iteration={n_iteration} | attempt={n_attempts} | {context}" - - try: - subprocess.run( - [ - sys.executable, - str(BASE_PATH / "notify.py"), - "--pipeline-id", - pipeline_id, - "--check-for", - "functional-tests", - "--pipeline-context", - pipeline_context, - "--pipeline-created-at", - pipeline_created_at, - ], - check=False, - ) - except Exception as e: - logger.warning("Failed to send Slack alert: %s", e) - - def register_pipeline_terminator(pipeline: jetclient.JETPipeline): def sigterm_handler(_signo, _stack_frame): print(f"Trying to terminate pipeline {pipeline.jet_id}") @@ -414,6 +375,8 @@ def is_flaky_failure(concat_allranks_logs: str) -> bool: or "free(): corrupted unsorted chunks" in concat_allranks_logs or "Segfault encountered" in concat_allranks_logs or "Fatal glibc error" in concat_allranks_logs + or "Disk quota exceeded" in concat_allranks_logs + or "basic_ios::clear: iostream error" in concat_allranks_logs ) @@ -644,33 +607,15 @@ def main( or "exiting program at iteration" in concat_allranks_logs ): logger.info("Release training finished") - send_slack_alert( - test_case=test_case, - context="training finished", - n_iteration=n_iteration, - n_attempts=n_attempts, - ) sys.exit(int(not success)) # invert for exit 0 if not success or parse_failed_job(logs=mainrank_log): logger.error("Release pipeline finished with status %s, retrying.", status.name) - send_slack_alert( - test_case=test_case, - context=f"pipeline finished with status {status.name}, retrying", - n_iteration=n_iteration, - n_attempts=n_attempts, - ) n_attempts += 1 continue n_iteration += 1 - send_slack_alert( - test_case=test_case, - context="max attempts exhausted", - n_iteration=n_iteration, - n_attempts=n_attempts, - ) telemetrics_and_exit( success=False, test_case=test_case, diff --git a/tests/test_utils/python_scripts/launch_nemo_run_workload.py b/tests/test_utils/python_scripts/launch_nemo_run_workload.py index 45b8086ea0a..b0d2f0a14fb 100644 --- a/tests/test_utils/python_scripts/launch_nemo_run_workload.py +++ b/tests/test_utils/python_scripts/launch_nemo_run_workload.py @@ -5,6 +5,7 @@ import os import pathlib import sys +import threading from typing import Optional import click @@ -23,6 +24,7 @@ def is_flaky_failure(concat_allranks_logs: str) -> bool: "The server socket has failed to listen on any local network address." in concat_allranks_logs or "Some NCCL operations have failed or timed out." in concat_allranks_logs + or "Watchdog caught collective operation timeout" in concat_allranks_logs or "uncorrectable ECC error encountered" in concat_allranks_logs or "illegal memory access" in concat_allranks_logs or "illegal instruction" in concat_allranks_logs @@ -53,6 +55,51 @@ def is_flaky_failure(concat_allranks_logs: str) -> bool: ) +def _is_hang_prone_flaky_failure(concat_allranks_logs: str) -> bool: + """Return whether a streamed failure may prevent the attempt from exiting.""" + return "Watchdog caught collective operation timeout" in concat_allranks_logs + + +class _ThreadSafeBuffer: + """Collect output shared between the log tailer and flaky-failure monitor.""" + + def __init__(self): + self._buffer = io.StringIO() + self._lock = threading.Lock() + + def write(self, data: str) -> None: + """Append log output to the buffer.""" + with self._lock: + self._buffer.write(data) + + def flush(self) -> None: + """Provide the stream interface expected by the tee wrapper.""" + + def getvalue(self) -> str: + """Return a consistent snapshot of the buffered output.""" + with self._lock: + return self._buffer.getvalue() + + +def _cancel_on_flaky_failure( + experiment: run.Experiment, + job_id: str, + log_buffer: _ThreadSafeBuffer, + stop_event: threading.Event, + failure_detected_event: threading.Event, + poll_interval: float = 1.0, +) -> None: + """Cancel an active attempt as soon as its streamed logs show a flaky failure.""" + while not stop_event.wait(poll_interval): + if _is_hang_prone_flaky_failure(log_buffer.getvalue()): + logger.warning( + "Detected flaky failure while job is running; cancelling current attempt." + ) + failure_detected_event.set() + experiment.cancel(job_id) + return + + def _collect_failure_logs(workdir: pathlib.Path) -> list[str]: """Reads every log file that may carry a flaky-failure signature. @@ -193,7 +240,7 @@ def main( n_attempts = 0 while n_attempts < 3: - tee_buffer = io.StringIO() + tee_buffer = _ThreadSafeBuffer() original_stdout = sys.stdout original_stderr = sys.stderr @@ -213,15 +260,33 @@ def flush(self): def __getattr__(self, name): return getattr(self._real, name) + monitor_stop_event = threading.Event() + flaky_failure_detected_event = threading.Event() + monitor_thread = None sys.stdout = _TeeStream(original_stdout, tee_buffer) sys.stderr = _TeeStream(original_stderr, tee_buffer) try: with run.Experiment("mcore-ci-test", executor=executor, log_level="INFO") as exp: - _ = exp.add([inline_script], tail_logs=False, name="task-1") + job_id = exp.add([inline_script], tail_logs=False, name="task-1") exp.dryrun(log=True) + monitor_thread = threading.Thread( + target=_cancel_on_flaky_failure, + args=( + exp, + job_id, + tee_buffer, + monitor_stop_event, + flaky_failure_detected_event, + ), + daemon=True, + ) + monitor_thread.start() exp.run(detach=False, tail_logs=True, sequential=False) finally: + monitor_stop_event.set() + if monitor_thread is not None: + monitor_thread.join() sys.stdout = original_stdout sys.stderr = original_stderr @@ -237,7 +302,7 @@ def __getattr__(self, name): all_ranks_all_logs = [tee_buffer.getvalue()] all_ranks_all_logs.extend(_collect_failure_logs(pathlib.Path(os.getcwd()))) all_ranks_all_logs_string = "\n".join(all_ranks_all_logs) - if is_flaky_failure(all_ranks_all_logs_string): + if flaky_failure_detected_event.is_set() or is_flaky_failure(all_ranks_all_logs_string): logger.warning("Detected flaky failure, attempt restart.") n_attempts += 1 continue diff --git a/tests/test_utils/python_scripts/linear_ci.py b/tests/test_utils/python_scripts/linear_ci.py new file mode 100644 index 00000000000..9bd9549b164 --- /dev/null +++ b/tests/test_utils/python_scripts/linear_ci.py @@ -0,0 +1,199 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Megatron-LM adapters for nemo-ci-triage's failure-reporting workflow. + +The triage package owns LLM summarization, Linear reconciliation, and Slack +follow-up logic. This module only converts Megatron-LM's direct child-pipeline +jobs into the generic failure records consumed by the package summarizer. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path +from typing import Any, Callable + +from nemo_ci_triage.agent import summarize_pipeline_failures as summarizer + +LINEAR_MODULE = "megatron_lm" +_FUNCTIONAL_PREFIX = "functional:run_" +_SMOKE_PREFIX = "functional:smoke-" + + +def _variant_name(pipeline_name: str) -> str: + """Return the stable environment/platform suffix of a functional bridge.""" + if pipeline_name.startswith(_FUNCTIONAL_PREFIX): + variant = pipeline_name.removeprefix(_FUNCTIONAL_PREFIX) + elif pipeline_name.startswith(_SMOKE_PREFIX): + variant = f"smoke-{pipeline_name.removeprefix(_SMOKE_PREFIX)}" + else: + variant = pipeline_name + return variant.replace("_", "-") + + +def _recipe_name(pipeline_name: str, config_name: str) -> str: + """Disambiguate the same recipe across dev/LTS and GPU child pipelines.""" + return f"{config_name}@{_variant_name(pipeline_name)}" + + +def _job_url(project_url: str, job: dict) -> str: + return job.get("web_url") or f"{project_url}/-/jobs/{job['id']}" + + +def _failure_record(pipeline_name: str, job: dict, report: dict | None, project_url: str) -> dict: + """Return the raw failure shape accepted by the upstream LLM summarizer.""" + return { + "test_name": _recipe_name(pipeline_name, job["config_name"]), + "module": LINEAR_MODULE, + "report": report, + "job_url": _job_url(project_url, job), + "job_error_type": job.get("error_type"), + } + + +def _fallback_summary(failure: dict) -> dict: + """Preserve a failed test when its per-test LLM summary is unavailable.""" + report = failure.get("report") or {} + category = ( + report.get("error_type") + or report.get("category") + or failure.get("job_error_type") + or "Unknown" + ) + subtype = report.get("error_subtype") or failure.get("job_error_type") + subtype = subtype or (f"No structured error report was available for {failure['test_name']}") + summary = subtype if subtype == category else f"{category}: {subtype}" + return { + "test_name": failure["test_name"], + "module": failure["module"], + "category": category, + "summary": summary, + "excerpt": report.get("excerpt"), + "job_url": failure["job_url"], + } + + +def _summarize_failures(raw_failures: list[dict]) -> list[dict]: + """Use upstream LLM summaries, falling back without dropping failures.""" + with_reports = [failure for failure in raw_failures if failure.get("report")] + summarized = summarizer._summarize_failures( + with_reports, summarizer._SUMMARIZER_PROMPT.read_text(encoding="utf-8").strip() + ) + by_job = {(failure["test_name"], failure["job_url"]): failure for failure in summarized} + return [ + by_job.get((failure["test_name"], failure["job_url"]), _fallback_summary(failure)) + for failure in raw_failures + ] + + +def build_pipeline_reports( + pipeline_id: int, + scope: str, + pipeline_jobs: list[tuple[str, int, list[dict]]], + load_error_report: Callable[[int], dict | None], + project_url: str, +) -> tuple[dict, dict]: + """Build the two JSON contracts consumed by nemo-ci-triage reconciliation. + + Each recipe is qualified by its child-pipeline variant. A recipe is only + included in ``passed_tests`` when that exact variant completed successfully; + failed, canceled, and ambiguous allow-failure jobs can therefore never close + a live Linear issue accidentally. + """ + passed: set[str] = set() + unknown: set[str] = set() + raw_failures: list[dict] = [] + failed_jobs = 0 + + for pipeline_name, _, jobs in sorted(pipeline_jobs, key=lambda item: item[0]): + for job in sorted(jobs, key=lambda item: (item["config_name"], item["id"])): + recipe = _recipe_name(pipeline_name, job["config_name"]) + status = job.get("status") + report = None + + if status == "failed" or (status == "success" and job.get("allow_failure")): + report = load_error_report(job["id"]) + + suppressed_failure = bool( + status == "success" and report and report.get("exit_code_training") not in (None, 0) + ) + if status == "failed" or suppressed_failure: + failed_jobs += 1 + raw_failures.append(_failure_record(pipeline_name, job, report, project_url)) + elif status == "success" and (not job.get("allow_failure") or report is not None): + passed.add(recipe) + else: + unknown.add(recipe) + + failed_recipes = {failure["test_name"] for failure in raw_failures} + passed_tests = sorted(passed - failed_recipes - unknown) + + failures = _summarize_failures(raw_failures) + buckets, failed_stage = summarizer._subcategorize(failures) + bucketing_failed = buckets is None + if bucketing_failed: + print( + f"WARNING: LLM categorizer failed at {failed_stage}; " + "Linear reconciliation will skip this report", + file=sys.stderr, + ) + buckets = [] + else: + summarizer._attach_categories(buckets, failures) + + digest = summarizer._digest( + failures, {LINEAR_MODULE: {"passed": len(passed_tests), "failed": failed_jobs}} + ) + + module_stats = { + "passed": len(passed_tests), + "failed": failed_jobs, + "passed_tests": passed_tests, + } + summaries = { + "pipeline_id": pipeline_id, + "scope": scope, + "modules": {LINEAR_MODULE: module_stats}, + "digest": digest, + "failures": failures, + } + failure_buckets = { + "pipeline_id": pipeline_id, + "bucketing_failed": bucketing_failed, + "buckets": summarizer._denormalize_buckets(buckets, failures), + } + return summaries, failure_buckets + + +def fetch_error_report(project: Any, job_id: int) -> dict | None: + """Fetch one child job's structured report, degrading safely if absent.""" + try: + raw = project.jobs.get(job_id, lazy=True).artifact("error_report.json") + if isinstance(raw, bytes): + raw = raw.decode("utf-8") + return json.loads(raw) + except Exception as exc: + print(f"WARNING: job {job_id}: could not read error_report.json: {exc}", file=sys.stderr) + return None + + +def write_pipeline_reports( + pipeline_id: int, + scope: str, + pipeline_jobs: list[tuple[str, int, list[dict]]], + project: Any, + project_url: str, + summaries_path: Path, + buckets_path: Path, +) -> None: + summaries, buckets = build_pipeline_reports( + pipeline_id, + scope, + pipeline_jobs, + lambda job_id: fetch_error_report(project, job_id), + project_url, + ) + summaries_path.write_text(json.dumps(summaries, indent=2) + "\n", encoding="utf-8") + buckets_path.write_text(json.dumps(buckets, indent=2) + "\n", encoding="utf-8") + print(f"Wrote {summaries_path} and {buckets_path}") diff --git a/tests/test_utils/python_scripts/notify.py b/tests/test_utils/python_scripts/notify.py index 103badc6ce5..71852bc617b 100644 --- a/tests/test_utils/python_scripts/notify.py +++ b/tests/test_utils/python_scripts/notify.py @@ -1,54 +1,90 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +import json import logging import os +from pathlib import Path +from typing import Any import click import gitlab -import pandas as pd -import requests -import slack_sdk +from nemo_ci_triage.slack_notification import notification +from nemo_ci_triage.slack_notification.utils import repository_settings -PROJECT_ID = int(os.getenv("CI_PROJECT_ID", 19378)) +from tests.test_utils.python_scripts import linear_ci + +TRIAGE_CONFIG = Path(os.getenv("NEMO_CI_TRIAGE_CONFIG", ".gitlab/nemo-ci-triage.yml")) +PROJECT_ID, REPO_NAME = repository_settings(TRIAGE_CONFIG) WEBHOOK_URL = os.getenv("WEBHOOK_URL", "") -GITLAB_ENDPOINT = os.getenv('GITLAB_ENDPOINT') -TAG_TEAM = bool(os.getenv('TAG_TEAM', 0)) -TEAM_SLUG = str(os.getenv('TEAM_SLUG')) +SLACK_BOT_TOKEN = os.getenv("MCORE_SLACK_BOT_TOKEN") or os.getenv("ALERTMANAGER_TOKEN", "") +SLACK_CHANNEL_ID = os.getenv("MCORE_SLACK_CHANNEL_ID", "") +GITLAB_ENDPOINT = os.getenv("GITLAB_ENDPOINT") +if not GITLAB_ENDPOINT: + raise ValueError("GITLAB_ENDPOINT is required") +SERVER_URL = f"https://{GITLAB_ENDPOINT}" +PROJECT_URL = os.getenv("CI_PROJECT_URL", f"{SERVER_URL}/{REPO_NAME}") +TAG_TEAM = os.getenv("TAG_TEAM", "0") == "1" +TEAM_SLUG = os.getenv("TEAM_SLUG", "") + +JOB_PREFIXES = { + "unit-tests": "test:unit_tests", + "integration-tests": "integration:run_", + "functional-tests": ("functional:run_", "functional:smoke-"), + "smoke-tests": "functional:smoke-", +} logging.basicConfig() logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) -def get_gitlab_handle(): - return gitlab.Gitlab(f"https://{GITLAB_ENDPOINT}", private_token=os.getenv("RO_API_TOKEN")) +def get_gitlab_handle() -> gitlab.Gitlab: + return gitlab.Gitlab(SERVER_URL, private_token=os.getenv("RO_API_TOKEN")) + + +def get_project() -> Any: + """Return the configured Megatron-LM GitLab project.""" + return get_gitlab_handle().projects.get(PROJECT_ID) + + +def _bridge_gpu(bridge_name: str) -> str: + for gpu in ("GB200", "H100", "A100"): + if gpu.lower() in bridge_name.lower(): + return gpu + return "Unknown" -def get_jobs_per_bridge(pipeline_id: int, type_of_job: str): - bridge = {} - for pipeline_bridge in ( - get_gitlab_handle() - .projects.get(PROJECT_ID) - .pipelines.get(pipeline_id) - .bridges.list(get_all=True) - ): - if ( - not pipeline_bridge.name.startswith(type_of_job) - or pipeline_bridge.attributes['downstream_pipeline'] is None - ): +def get_pipeline_jobs( + pipeline_id: int, job_prefix: str | tuple[str, ...], project: Any | None = None +) -> list[tuple[str, int, list[dict]]]: + """Collect Megatron-LM's direct child pipelines using nemo-ci-triage-2.""" + project = project or get_project() + root_pipeline = project.pipelines.get(pipeline_id) + pipeline_jobs = [] + + for bridge in root_pipeline.bridges.list(get_all=True): + downstream = bridge.attributes.get("downstream_pipeline") + if not bridge.name.startswith(job_prefix) or downstream is None: continue - if pipeline_bridge.name not in bridge: - bridge[pipeline_bridge.name] = [] + child_pipeline_id = downstream["id"] + jobs = notification.get_jobs_from_pipeline(project, child_pipeline_id) + bridge_gpu = _bridge_gpu(bridge.name) + for job in jobs: + if job["gpu"] == "Unknown": + job["gpu"] = bridge_gpu + pipeline_jobs.append((bridge.name, child_pipeline_id, jobs)) + + return pipeline_jobs - for job in ( - get_gitlab_handle() - .projects.get(PROJECT_ID) - .pipelines.get(pipeline_bridge.attributes['downstream_pipeline']['id']) - .jobs.list(get_all=True) - ): - bridge[pipeline_bridge.name].append(job) - return bridge + +def write_slack_context(output: Path | None, thread_timestamp: str | None) -> None: + """Persist the non-secret Slack coordinates needed by a follow-up job.""" + if output is None: + return + context = {"channel_id": SLACK_CHANNEL_ID or None, "thread_timestamp": thread_timestamp} + output.write_text(json.dumps(context, indent=2) + "\n", encoding="utf-8") + logger.info("Wrote Slack thread context to %s", output) @click.command() @@ -59,59 +95,67 @@ def get_jobs_per_bridge(pipeline_id: int, type_of_job: str): type=click.Choice(["unit-tests", "integration-tests", "functional-tests", "smoke-tests"]), ) @click.option("--pipeline-context", required=True, type=str) -@click.option("--pipeline-created-at", required=True, type=str) -def main(pipeline_id: int, check_for: str, pipeline_context: str, pipeline_created_at: str): - if check_for == "unit-tests": - bridges = get_jobs_per_bridge(pipeline_id, "test:unit_tests") - - if check_for == "integration-tests": - bridges = get_jobs_per_bridge(pipeline_id, "integration:run_") +@click.option("--pipeline-created-at", required=True, type=str, expose_value=False) +@click.option("--summary-output", type=click.Path(path_type=Path), default=None) +@click.option("--failure-buckets-output", type=click.Path(path_type=Path), default=None) +@click.option("--slack-output", type=click.Path(path_type=Path), default=None) +def main( + pipeline_id: int, + check_for: str, + pipeline_context: str, + summary_output: Path | None, + failure_buckets_output: Path | None, + slack_output: Path | None, +) -> None: + if bool(summary_output) != bool(failure_buckets_output): + raise click.UsageError( + "--summary-output and --failure-buckets-output must be provided together" + ) - if check_for == "functional-tests": - bridges = get_jobs_per_bridge(pipeline_id, "functional:run_") + project = get_project() + pipeline_jobs = get_pipeline_jobs(pipeline_id, JOB_PREFIXES[check_for], project=project) + + if summary_output: + linear_ci.write_pipeline_reports( + pipeline_id, + pipeline_context, + pipeline_jobs, + project, + PROJECT_URL, + summary_output, + failure_buckets_output, + ) if check_for == "smoke-tests": - bridges = get_jobs_per_bridge(pipeline_id, "functional:smoke-") - if all(job.status == "success" for jobs in bridges.values() for job in jobs): + if all(job["status"] == "success" for _, _, jobs in pipeline_jobs for job in jobs): logger.info("All smoke tests passed, skipping Slack notification") + write_slack_context(slack_output, None) return - pipeline_created_at_day = pd.Timestamp(pipeline_created_at).strftime("%Y-%m-%d") - - messages = [] - - for bridge_name in bridges.keys(): - - total_num_jobs = len(bridges[bridge_name]) - if all(job.status == "success" for job in bridges[bridge_name]): - messages.append( - f":doge3d: : All {total_num_jobs} passed." - ) - continue - - unsuccessful_jobs = [job for job in bridges[bridge_name] if job.status != "success"] - messages.append( - f":doctorge: : {len(unsuccessful_jobs)} of {total_num_jobs} failed." + use_bot = bool(SLACK_BOT_TOKEN and SLACK_CHANNEL_ID) + if bool(SLACK_BOT_TOKEN) != bool(SLACK_CHANNEL_ID): + logger.warning( + "Both MCORE_SLACK_BOT_TOKEN (or ALERTMANAGER_TOKEN) and " + "MCORE_SLACK_CHANNEL_ID are required for threaded Slack replies" ) - if TAG_TEAM: - messages.append( - f"cc {TEAM_SLUG} <@U09TX0DHZ97>: Critical event, please react as soon as possible." - ) - - for job in unsuccessful_jobs: - messages.append( - f"\tJob: " - ) - - messages.append("===============================================") - if not WEBHOOK_URL: - logger.info("No webhook URL configured, skipping Slack notification") + if not WEBHOOK_URL and not use_bot: + logger.info("No Slack bot or webhook configured, skipping Slack notification") + write_slack_context(slack_output, None) return - for message in messages: - response = slack_sdk.webhook.WebhookClient(WEBHOOK_URL).send(text=message) - logger.info(response.status_code) + slack_mentions = f"{TEAM_SLUG} <@U09TX0DHZ97>" if TAG_TEAM else None + thread_timestamp = notification.send_slack_notification( + "megatron-lm", + pipeline_context, + pipeline_jobs, + slack_mentions, + webhook_url=WEBHOOK_URL or None, + slack_bot_token=SLACK_BOT_TOKEN if use_bot else None, + slack_channel_id=SLACK_CHANNEL_ID if use_bot else None, + config=TRIAGE_CONFIG, + ) + write_slack_context(slack_output, thread_timestamp) if __name__ == "__main__": diff --git a/tests/test_utils/recipes/gb200/gpt-dynamic-inference.yaml b/tests/test_utils/recipes/gb200/gpt-dynamic-inference.yaml new file mode 100644 index 00000000000..6a81fab3b9a --- /dev/null +++ b/tests/test_utils/recipes/gb200/gpt-dynamic-inference.yaml @@ -0,0 +1,65 @@ +type: basic +format_version: 1 +maintainers: [mcore] +loggers: [stdout] +spec: + name: '{test_case}_{environment}_{platforms}' + model: gpt + build: mcore-pyt-{environment} + nodes: 1 + gpus: 4 + n_repeat: 1 + platforms: dgx_gb200 + script_setup: | + set -euo pipefail + unset https_proxy + echo "machine gitlab-master.nvidia.com login okoenig password $RO_API_TOKEN" | tee -a /root/.netrc + + # Checkout latest + cd /opt + rm -rf /opt/megatron-lm; mkdir megatron-lm; cd megatron-lm + git init + git remote add origin $MCORE_REPO + git fetch origin '+refs/merge-requests/*:refs/remotes/merge-requests/*' + git fetch origin $MCORE_MR_COMMIT + git checkout $MCORE_MR_COMMIT + git rev-parse HEAD + # Checkout backwards-ref + cd /opt + rm -rf /opt/megatron-lm-legacy; mkdir megatron-lm-legacy; cd megatron-lm-legacy + git init + git remote add origin $MCORE_REPO + git fetch origin $MCORE_BACKWARDS_COMMIT + git checkout $MCORE_BACKWARDS_COMMIT + git rev-parse HEAD + rm -rf megatron; cp -a /opt/megatron-lm/megatron ./ + script: |- + set -euo pipefail + ls + cd /opt/megatron-lm + export GPUS_PER_NODE={gpus} + + ARGUMENTS=( + "CHECKPOINT_LOAD_PATH=/mnt/artifacts" + "CHECKPOINT_SAVE_PATH=/tmp/checkpoints" + "DATA_PATH=/mnt/artifacts" + "DATA_CACHE_PATH=/workspace/data/cache" + "TRAINING_SCRIPT_PATH=examples/inference/advanced/gpt_dynamic_inference.py" + "TRAINING_PARAMS_PATH=./tests/functional_tests/test_cases/{model}/{test_case}/model_config.yaml" + "GOLDEN_VALUES_PATH=./tests/functional_tests/test_cases/{model}/{test_case}/golden_values_{environment}_{platforms}.json" + "OUTPUT_PATH={assets_dir}" + "TENSORBOARD_PATH={assets_dir}/tensorboard" + "INFERENCE_OUTPUT_PATH={assets_dir}/golden_values_{environment}_{platforms}.json" + "N_REPEAT={n_repeat}" + "ENABLE_LIGHTWEIGHT_MODE=${{ENABLE_LIGHTWEIGHT_MODE:-}}" + "RECORD_CHECKPOINTS=${{RECORD_CHECKPOINTS:-}}" + ) + + bash ./tests/functional_tests/shell_test_utils/run_ci_test.sh ${{ARGUMENTS[@]}} + +products: + - test_case: [gpt_dynamic_inference_tp1_pp1_583m_async_sched] + products: + - environment: [dev] + scope: [mr] + platforms: [dgx_gb200] diff --git a/tests/test_utils/recipes/gb200/gpt-perf-dp4.yaml b/tests/test_utils/recipes/gb200/gpt-perf-dp4.yaml index 9f4839d28fe..71a9dccf1b2 100644 --- a/tests/test_utils/recipes/gb200/gpt-perf-dp4.yaml +++ b/tests/test_utils/recipes/gb200/gpt-perf-dp4.yaml @@ -40,7 +40,7 @@ spec: GPUS_PER_NODE=4 bash ./tests/performance_tests/shell_test_utils/run_perf_test.sh ${{ARGUMENTS[@]}} products: - - test_case: [gpt_583m_perf_gb200_4gpu] + - test_case: [gpt_583m_perf_gb200_4gpu, gpt_583m_perf_async_sched_gb200_4gpu] products: - environment: [dev] scope: [mr] diff --git a/tests/test_utils/recipes/h100/gpt-dynamic-inference-with-coordinator.yaml b/tests/test_utils/recipes/h100/gpt-dynamic-inference-with-coordinator.yaml index fdc96221e44..1cc8c1a47ae 100644 --- a/tests/test_utils/recipes/h100/gpt-dynamic-inference-with-coordinator.yaml +++ b/tests/test_utils/recipes/h100/gpt-dynamic-inference-with-coordinator.yaml @@ -82,7 +82,7 @@ products: - environment: [dev] scope: [mr] platforms: [dgx_h100] - - test_case: [gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_round_robin_zmq] + - test_case: [gpt_dynamic_inference_tp1_pp1_dp8_583m_prefix_caching_load_balanced_zmq] products: - environment: [dev] scope: [mr] diff --git a/tests/test_utils/recipes/h100/gpt-dynamic-inference.yaml b/tests/test_utils/recipes/h100/gpt-dynamic-inference.yaml index 43661c16cd3..79e97beb5f0 100644 --- a/tests/test_utils/recipes/h100/gpt-dynamic-inference.yaml +++ b/tests/test_utils/recipes/h100/gpt-dynamic-inference.yaml @@ -178,3 +178,8 @@ products: - environment: [dev] scope: [mr] platforms: [dgx_h100] + - test_case: [gpt_dynamic_inference_tp1_pp1_583m_async_sched] + products: + - environment: [dev] + scope: [mr] + platforms: [dgx_h100] diff --git a/tests/test_utils/recipes/h100/gpt-perf-dp8.yaml b/tests/test_utils/recipes/h100/gpt-perf-dp8.yaml index 6d9cc4e948f..7484989358f 100644 --- a/tests/test_utils/recipes/h100/gpt-perf-dp8.yaml +++ b/tests/test_utils/recipes/h100/gpt-perf-dp8.yaml @@ -38,7 +38,7 @@ spec: GPUS_PER_NODE=8 bash ./tests/performance_tests/shell_test_utils/run_perf_test.sh ${{ARGUMENTS[@]}} products: - - test_case: [gpt_583m_perf] + - test_case: [gpt_583m_perf, gpt_583m_perf_async_sched] products: - environment: [dev] scope: [mr] diff --git a/tests/test_utils/recipes/h100/gpt.yaml b/tests/test_utils/recipes/h100/gpt.yaml index 3c808d3466b..12c67461ea4 100644 --- a/tests/test_utils/recipes/h100/gpt.yaml +++ b/tests/test_utils/recipes/h100/gpt.yaml @@ -197,8 +197,12 @@ products: scope: [nightly] - test_case: [gpt3_mcore_te_tp1_pp4_vp1_resume_torch_decoupled_lr] products: + # Disabled (flaky): the exact/deterministic golden comparison intermittently + # fails on dgx_h100 while the approximate check passes and the run is + # reproducible in isolation. Kept in-tree for easy re-enable; see the + # existing "WAR to #513" note on this test case's model_config.yaml. - environment: [dev] - scope: [mr, mr-github] + scope: [mr-broken, mr-github-broken] platforms: [dgx_h100] - environment: [lts] scope: [nightly] diff --git a/tests/test_utils/recipes/h100/mamba.yaml b/tests/test_utils/recipes/h100/mamba.yaml index 4ec2f0d62ea..2675ecbdc7c 100644 --- a/tests/test_utils/recipes/h100/mamba.yaml +++ b/tests/test_utils/recipes/h100/mamba.yaml @@ -102,3 +102,10 @@ products: platforms: [dgx_h100] # - environment: [lts] # disabled until triton is bumped # scope: [nightly] + + - test_case: [hybrid_nemotron_v3_pico_7b_a1b_tp1_ep8_QAD_dgx_h100_1N8G] + products: + # Disabled while the deterministic total-loss mismatch is investigated. + - environment: [dev] + scope: [mr-broken, mr-github-broken] + platforms: [dgx_h100] diff --git a/tests/test_utils/test_ci_triage.py b/tests/test_utils/test_ci_triage.py new file mode 100644 index 00000000000..b6494c66ea0 --- /dev/null +++ b/tests/test_utils/test_ci_triage.py @@ -0,0 +1,580 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import json +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest +import yaml +from click.testing import CliRunner + +from tests.test_utils.python_scripts import generate_jet_trigger_job, linear_ci, recipe_parser + + +def _mock_llm_reporting(monkeypatch): + summarize = Mock( + side_effect=lambda failures, _prompt: [ + linear_ci._fallback_summary(failure) for failure in failures + ] + ) + + def group_failures(failures): + grouped = {} + for failure in failures: + grouped.setdefault((failure["category"], failure["summary"]), []).append( + failure["test_name"] + ) + return ( + [ + {"label": f"test-bucket-{index}", "rationale": summary, "tests": tests} + for index, ((_, summary), tests) in enumerate(grouped.items(), 1) + ], + None, + ) + + subcategorize = Mock(side_effect=group_failures) + digest = Mock(return_value="LLM pipeline digest") + monkeypatch.setattr(linear_ci.summarizer, "_summarize_failures", summarize) + monkeypatch.setattr(linear_ci.summarizer, "_subcategorize", subcategorize) + monkeypatch.setattr(linear_ci.summarizer, "_digest", digest) + return summarize, subcategorize, digest + + +@pytest.fixture +def notify_module(monkeypatch): + pytest.importorskip("nemo_ci_triage.slack_notification") + monkeypatch.setenv("GITLAB_ENDPOINT", "ci.example.com") + from tests.test_utils.python_scripts import notify + + return notify + + +def test_build_test_script_preserves_workload_exit_code(): + script = generate_jet_trigger_job.build_test_script("python workload.py") + + assert "set +e" in script + assert "set -o pipefail" in script + assert "python workload.py 2>&1 | tee jet_workload.log" in script + assert "exit_code=${PIPESTATUS[0]}" in script + assert "set -e" in script + assert "extract-errors jet_workload.log" in script + assert "--output error_report.json" in script + assert 'exit "$exit_code"' in script + + +@pytest.mark.parametrize("enable_error_extraction", [False, True]) +def test_error_extraction_is_opt_in_for_generated_jobs( + monkeypatch, tmp_path, enable_error_extraction +): + workload = recipe_parser.dotdict( + type="basic", + spec=recipe_parser.dotdict(model="gpt", environment="dev", test_case="triage-test"), + ) + monkeypatch.setattr( + generate_jet_trigger_job.recipe_parser, "load_workloads", lambda **_kwargs: [workload] + ) + output_path = tmp_path / "pipeline.yml" + args = [ + "--scope", + "mr", + "--environment", + "dev", + "--n-repeat", + "1", + "--time-limit", + "60", + "--test-cases", + "all", + "--platform", + "dgx_h100", + "--cluster", + "ghci", + "--output-path", + str(output_path), + "--container-image", + "utility", + "--container-tag", + "test", + "--dependent-job", + "functional:configure", + "--record-checkpoints", + "false", + "--slurm-account", + "mcore", + "--no-enable-warmup", + ] + if enable_error_extraction: + args.append("--enable-error-extraction") + + result = CliRunner().invoke(generate_jet_trigger_job.main, args) + + assert result.exit_code == 0, result.output + job = yaml.safe_load(output_path.read_text())["triage-test"] + if enable_error_extraction: + assert "extract-errors jet_workload.log" in job["script"][0] + assert job["artifacts"]["paths"] == ["results/", "jet_workload.log", "error_report.json"] + else: + assert "extract-errors" not in job["script"][0] + assert job["artifacts"]["paths"] == ["results/"] + + +def test_notification_rules_use_expected_pipeline_sources(): + unit = yaml.safe_load(Path(".gitlab/stages/02.test.yml").read_text()) + functional = yaml.safe_load(Path(".gitlab/stages/04.functional-tests.yml").read_text()) + triage = yaml.safe_load(Path(".gitlab/stages/06.triage.yml").read_text()) + + unit_conditions = [ + rule["if"] for rule in unit["test:unit_tests_notify"]["rules"] if "if" in rule + ] + assert unit_conditions == [ + '$CI_PIPELINE_SOURCE == "schedule" && ' + '($CI_COMMIT_BRANCH == "ci-unit-test-extended" || ' + '$CI_COMMIT_BRANCH == "ci-dev-unit-test-extended")' + ] + + assert "functional:smoke_notify" not in functional + assert functional["functional:x_notify"]["rules"][0]["if"] == ( + '($CI_PIPELINE_SOURCE == "schedule" || $CI_COMMIT_BRANCH == "main") && ' + '$FUNCTIONAL_TEST == "yes"' + ) + + triage_jobs = (".linear_reconcile_rules", "triage:linear_write", "triage:slack_linear_followup") + for job_name in triage_jobs: + condition = triage[job_name]["rules"][0]["if"] + assert '$FUNCTIONAL_TEST == "yes"' in condition + assert '$CI_PIPELINE_SOURCE == "schedule"' in condition + assert '$CI_COMMIT_BRANCH == "main"' in condition + + +def test_all_generated_test_types_enable_error_extraction(): + unit = Path(".gitlab/stages/02.test.yml").read_text() + integration = Path(".gitlab/stages/03.integration-tests.yml").read_text() + functional = Path(".gitlab/stages/04.functional-tests.yml").read_text() + + assert unit.count('"--enable-error-extraction"') >= 1 + assert integration.count('"--enable-error-extraction"') >= 1 + assert functional.count('"--enable-error-extraction"') >= 2 + + +def test_functional_notifications_are_parent_aggregate_only(): + launcher = Path("tests/test_utils/python_scripts/launch_jet_workload.py").read_text() + functional = yaml.safe_load(Path(".gitlab/stages/04.functional-tests.yml").read_text()) + notify_script = "\n".join(functional["functional:x_notify"]["script"]) + + assert "send_slack_alert" not in launcher + assert "notify.py" not in launcher + assert notify_script.count("python tests/test_utils/python_scripts/notify.py") == 1 + assert "--check-for functional-tests" in notify_script + assert "--pipeline-context $CONTEXT" in notify_script + + +def test_get_pipeline_jobs_uses_triage_collector(monkeypatch, notify_module): + notify = notify_module + bridges = [ + SimpleNamespace( + name="functional:run_dev_dgx_h100", attributes={"downstream_pipeline": {"id": 101}} + ), + SimpleNamespace( + name="functional:smoke-gb200", attributes={"downstream_pipeline": {"id": 102}} + ), + ] + root_pipeline = Mock() + root_pipeline.bridges.list.return_value = bridges + project = Mock() + project.pipelines.get.return_value = root_pipeline + handle = Mock() + handle.projects.get.return_value = project + monkeypatch.setattr(notify, "get_gitlab_handle", lambda: handle) + collector = Mock( + side_effect=[ + [{"status": "failed", "gpu": "Unknown"}], + [{"status": "failed", "gpu": "Unknown"}], + ] + ) + monkeypatch.setattr(notify.notification, "get_jobs_from_pipeline", collector) + + assert notify.get_pipeline_jobs(123, notify.JOB_PREFIXES["functional-tests"]) == [ + ("functional:run_dev_dgx_h100", 101, [{"status": "failed", "gpu": "H100"}]), + ("functional:smoke-gb200", 102, [{"status": "failed", "gpu": "GB200"}]), + ] + assert collector.call_args_list == [((project, 101),), ((project, 102),)] + + +def test_build_linear_reports_groups_matching_failures(monkeypatch): + summarize, subcategorize, digest = _mock_llm_reporting(monkeypatch) + pipeline_jobs = [ + ( + "functional:run_dev_dgx_h100", + 101, + [ + { + "config_name": "gpt_pass", + "id": 1, + "status": "success", + "allow_failure": False, + "error_type": None, + }, + { + "config_name": "gpt_fail_a", + "id": 2, + "status": "failed", + "allow_failure": False, + "error_type": "CUDA OOM", + }, + ], + ), + ( + "functional:run_lts_dgx_h100", + 102, + [ + { + "config_name": "gpt_fail_b", + "id": 3, + "status": "failed", + "allow_failure": True, + "error_type": None, + } + ], + ), + ( + "functional:smoke-gb200", + 103, + [ + { + "config_name": "gpt_smoke_fail", + "id": 4, + "status": "failed", + "allow_failure": False, + "error_type": "CUDA OOM", + } + ], + ), + ] + reports = { + 2: { + "exit_code_training": 1, + "category": "CUDA OOM", + "error_subtype": "torch.OutOfMemoryError", + "excerpt": "CUDA out of memory", + }, + 3: { + "exit_code_training": 1, + "category": "CUDA OOM", + "error_subtype": "torch.OutOfMemoryError", + "excerpt": "CUDA out of memory", + }, + 4: { + "exit_code_training": 1, + "category": "CUDA OOM", + "error_subtype": "torch.OutOfMemoryError", + "excerpt": "CUDA out of memory", + }, + } + + summaries, buckets = linear_ci.build_pipeline_reports( + 123, "nightly", pipeline_jobs, reports.get, "https://ci.example.com/ADLR/megatron-lm" + ) + + stats = summaries["modules"][linear_ci.LINEAR_MODULE] + assert stats == {"passed": 1, "failed": 3, "passed_tests": ["gpt_pass@dev-dgx-h100"]} + assert len(buckets["buckets"]) == 1 + bucket = buckets["buckets"][0] + assert bucket["module"] == linear_ci.LINEAR_MODULE + assert bucket["category"] == "CUDA OOM" + assert bucket["rationale"] == "CUDA OOM: torch.OutOfMemoryError" + assert bucket["tests"] == [ + { + "name": "gpt_fail_a@dev-dgx-h100", + "job_url": "https://ci.example.com/ADLR/megatron-lm/-/jobs/2", + }, + { + "name": "gpt_fail_b@lts-dgx-h100", + "job_url": "https://ci.example.com/ADLR/megatron-lm/-/jobs/3", + }, + { + "name": "gpt_smoke_fail@smoke-gb200", + "job_url": "https://ci.example.com/ADLR/megatron-lm/-/jobs/4", + }, + ] + summarize.assert_called_once() + subcategorize.assert_called_once() + digest.assert_called_once() + + +def test_allow_failure_without_report_is_not_counted_as_passed(monkeypatch): + _mock_llm_reporting(monkeypatch) + pipeline_jobs = [ + ( + "functional:run_dev_dgx_h100", + 101, + [ + { + "config_name": "ambiguous", + "id": 4, + "status": "success", + "allow_failure": True, + "error_type": None, + } + ], + ) + ] + + summaries, buckets = linear_ci.build_pipeline_reports( + 123, + "nightly", + pipeline_jobs, + lambda _job_id: None, + "https://ci.example.com/ADLR/megatron-lm", + ) + + stats = summaries["modules"][linear_ci.LINEAR_MODULE] + assert stats["passed_tests"] == [] + assert stats["failed"] == 0 + assert buckets["buckets"] == [] + + +def test_failed_job_without_report_still_creates_a_safe_bucket(monkeypatch): + _mock_llm_reporting(monkeypatch) + pipeline_jobs = [ + ( + "functional:run_dev_dgx_h100", + 101, + [ + { + "config_name": "missing_report", + "id": 5, + "status": "failed", + "allow_failure": False, + "error_type": None, + } + ], + ) + ] + + summaries, buckets = linear_ci.build_pipeline_reports( + 123, + "nightly", + pipeline_jobs, + lambda _job_id: None, + "https://ci.example.com/ADLR/megatron-lm", + ) + + assert summaries["modules"][linear_ci.LINEAR_MODULE]["failed"] == 1 + assert buckets["buckets"][0]["tests"][0]["name"] == "missing_report@dev-dgx-h100" + assert "No structured error report" in buckets["buckets"][0]["rationale"] + + +def test_triage_config_selects_megatron_and_enables_write_actions(): + linear_status = pytest.importorskip("nemo_ci_triage.linear.linear_status") + linear_write = pytest.importorskip("nemo_ci_triage.linear.linear_write") + config = Path(".gitlab/nemo-ci-triage.yml") + + assert linear_status.modules_for_regex("^megatron-lm$", config) == [ + ( + linear_ci.LINEAR_MODULE, + { + "build_module": "megatron-lm", + "channel_id_env": "MCORE_SLACK_CHANNEL_ID", + "reconcile_proposal": True, + "team_key": "MCORE", + "project_template": "MCore CI Testing", + "enable_linear_open": True, + "enable_linear_modify": True, + "enable_linear_close": True, + }, + ) + ] + assert linear_write.write_gates(config) == { + linear_ci.LINEAR_MODULE: {"open": True, "modify": True, "close": True} + } + + +def test_slack_followup_uses_upstream_detailed_and_execution_summaries(): + triage = yaml.safe_load(Path(".gitlab/stages/06.triage.yml").read_text()) + execution_summary, detailed_summary = triage["triage:slack_linear_followup"]["script"] + script = f"{execution_summary}\n{detailed_summary}" + + assert "--pipeline-summary slack_notification.json" in execution_summary + assert "--linear-plan linear_action_plan_post.json" in execution_summary + assert "--slack-channel-id" in execution_summary + assert "--module megatron_lm" in detailed_summary + assert "--only-followup" in detailed_summary + assert '--thread-ts "${THREAD_TIMESTAMP}"' in detailed_summary + assert "--failure-buckets failure_buckets.json" in detailed_summary + assert "--linear-report linear_status_report.json" in detailed_summary + assert "--action-plan linear_action_plan_post.json" in detailed_summary + assert "--slack-channel-id" not in detailed_summary + assert 'if [[ -z "${THREAD_TIMESTAMP}" ]]' in script + + +@pytest.mark.parametrize("pipeline_context", ["mr", "nightly", "weekly", "release"]) +def test_notification_delegates_to_triage_package(monkeypatch, notify_module, pipeline_context): + notify = notify_module + project = Mock() + pipeline_jobs = [("functional:run_dev_dgx_h100", 101, [{"status": "failed"}])] + collector = Mock(return_value=pipeline_jobs) + sender = Mock() + + monkeypatch.setattr(notify, "WEBHOOK_URL", "https://slack.invalid/webhook") + monkeypatch.setattr(notify, "SLACK_BOT_TOKEN", "") + monkeypatch.setattr(notify, "SLACK_CHANNEL_ID", "") + monkeypatch.setattr(notify, "PROJECT_URL", "https://ci.example.com/ADLR/megatron-lm") + monkeypatch.setattr(notify, "get_project", lambda: project) + monkeypatch.setattr(notify, "get_pipeline_jobs", collector) + monkeypatch.setattr(notify.notification, "send_slack_notification", sender) + + result = CliRunner().invoke( + notify.main, + [ + "--pipeline-id", + "123", + "--check-for", + "functional-tests", + "--pipeline-context", + pipeline_context, + "--pipeline-created-at", + "2026-07-12T00:00:00Z", + ], + ) + + assert result.exit_code == 0, result.output + sender.assert_called_once_with( + "megatron-lm", + pipeline_context, + pipeline_jobs, + None, + webhook_url="https://slack.invalid/webhook", + slack_bot_token=None, + slack_channel_id=None, + config=notify.TRIAGE_CONFIG, + ) + collector.assert_called_once_with(123, notify.JOB_PREFIXES["functional-tests"], project=project) + + +@pytest.mark.parametrize("has_failure", [False, True]) +def test_smoke_notification_is_failure_only_and_aggregate(monkeypatch, notify_module, has_failure): + notify = notify_module + project = Mock() + pipeline_jobs = [ + ("functional:smoke-h100", 101, [{"status": "success"}]), + ("functional:smoke-gb200", 102, [{"status": "failed" if has_failure else "success"}]), + ] + collector = Mock(return_value=pipeline_jobs) + sender = Mock() + + monkeypatch.setattr(notify, "WEBHOOK_URL", "https://slack.invalid/webhook") + monkeypatch.setattr(notify, "SLACK_BOT_TOKEN", "") + monkeypatch.setattr(notify, "SLACK_CHANNEL_ID", "") + monkeypatch.setattr(notify, "get_project", lambda: project) + monkeypatch.setattr(notify, "get_pipeline_jobs", collector) + monkeypatch.setattr(notify.notification, "send_slack_notification", sender) + + result = CliRunner().invoke( + notify.main, + [ + "--pipeline-id", + "123", + "--check-for", + "smoke-tests", + "--pipeline-context", + "smoke-nightly", + "--pipeline-created-at", + "2026-07-12T00:00:00Z", + ], + ) + + assert result.exit_code == 0, result.output + collector.assert_called_once_with(123, notify.JOB_PREFIXES["smoke-tests"], project=project) + if has_failure: + sender.assert_called_once() + assert sender.call_args.args[2] == pipeline_jobs + else: + sender.assert_not_called() + + +def test_notification_records_bot_thread_context(monkeypatch, tmp_path, notify_module): + notify = notify_module + pipeline_jobs = [("functional:run_dev_dgx_h100", 101, [{"status": "failed"}])] + sender = Mock(return_value="1712345678.000100") + slack_output = tmp_path / "slack_notification.json" + + monkeypatch.setattr(notify, "WEBHOOK_URL", "") + monkeypatch.setattr(notify, "SLACK_BOT_TOKEN", "xoxb-test") + monkeypatch.setattr(notify, "SLACK_CHANNEL_ID", "C0123456789") + monkeypatch.setattr(notify, "get_project", Mock()) + monkeypatch.setattr(notify, "get_pipeline_jobs", lambda *_args, **_kwargs: pipeline_jobs) + monkeypatch.setattr(notify.notification, "send_slack_notification", sender) + + result = CliRunner().invoke( + notify.main, + [ + "--pipeline-id", + "123", + "--check-for", + "functional-tests", + "--pipeline-context", + "mr", + "--pipeline-created-at", + "2026-07-12T00:00:00Z", + "--slack-output", + str(slack_output), + ], + ) + + assert result.exit_code == 0, result.output + sender.assert_called_once_with( + "megatron-lm", + "mr", + pipeline_jobs, + None, + webhook_url=None, + slack_bot_token="xoxb-test", + slack_channel_id="C0123456789", + config=notify.TRIAGE_CONFIG, + ) + assert json.loads(slack_output.read_text()) == { + "channel_id": "C0123456789", + "thread_timestamp": "1712345678.000100", + } + + +def test_notification_writes_linear_inputs_without_webhook(monkeypatch, tmp_path, notify_module): + notify = notify_module + project = Mock() + pipeline_jobs = [("functional:run_dev_dgx_h100", 101, [])] + collector = Mock(return_value=pipeline_jobs) + writer = Mock() + + monkeypatch.setattr(notify, "WEBHOOK_URL", "") + monkeypatch.setattr(notify, "SLACK_BOT_TOKEN", "") + monkeypatch.setattr(notify, "SLACK_CHANNEL_ID", "") + monkeypatch.setattr(notify, "get_project", lambda: project) + monkeypatch.setattr(notify, "get_pipeline_jobs", collector) + monkeypatch.setattr(notify.linear_ci, "write_pipeline_reports", writer) + summaries = tmp_path / "pipeline_summaries.json" + buckets = tmp_path / "failure_buckets.json" + + result = CliRunner().invoke( + notify.main, + [ + "--pipeline-id", + "123", + "--check-for", + "functional-tests", + "--pipeline-context", + "nightly", + "--pipeline-created-at", + "2026-07-12T00:00:00Z", + "--summary-output", + str(summaries), + "--failure-buckets-output", + str(buckets), + ], + ) + + assert result.exit_code == 0, result.output + collector.assert_called_once_with(123, notify.JOB_PREFIXES["functional-tests"], project=project) + writer.assert_called_once_with( + 123, "nightly", pipeline_jobs, project, notify.PROJECT_URL, summaries, buckets + ) diff --git a/tests/test_utils/test_launch_nemo_run_workload.py b/tests/test_utils/test_launch_nemo_run_workload.py new file mode 100644 index 00000000000..ea9930b1fa8 --- /dev/null +++ b/tests/test_utils/test_launch_nemo_run_workload.py @@ -0,0 +1,68 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import threading +from unittest.mock import Mock + +from tests.test_utils.python_scripts import launch_nemo_run_workload + + +def test_nccl_watchdog_timeout_is_flaky(): + log = "Watchdog caught collective operation timeout: WorkNCCL(SeqNum=281)" + + assert launch_nemo_run_workload.is_flaky_failure(log) + + +def test_hang_prone_flaky_failure_cancels_active_attempt(): + experiment = Mock() + log_buffer = launch_nemo_run_workload._ThreadSafeBuffer() + stop_event = threading.Event() + failure_detected_event = threading.Event() + monitor = threading.Thread( + target=launch_nemo_run_workload._cancel_on_flaky_failure, + args=(experiment, "task-1", log_buffer, stop_event, failure_detected_event, 0.01), + ) + + monitor.start() + log_buffer.write("Watchdog caught collective operation timeout") + monitor.join(timeout=1) + + assert not monitor.is_alive() + assert failure_detected_event.is_set() + experiment.cancel.assert_called_once_with("task-1") + + +def test_non_hanging_flaky_failure_does_not_cancel_active_attempt(): + experiment = Mock() + log_buffer = launch_nemo_run_workload._ThreadSafeBuffer() + log_buffer.write("found NaN in local forward loss calculation") + stop_event = threading.Event() + failure_detected_event = threading.Event() + monitor = threading.Thread( + target=launch_nemo_run_workload._cancel_on_flaky_failure, + args=(experiment, "task-1", log_buffer, stop_event, failure_detected_event, 0.01), + ) + + monitor.start() + assert not failure_detected_event.wait(timeout=0.05) + stop_event.set() + monitor.join(timeout=1) + + assert not monitor.is_alive() + assert launch_nemo_run_workload.is_flaky_failure(log_buffer.getvalue()) + experiment.cancel.assert_not_called() + + +def test_stopped_monitor_does_not_cancel_attempt(): + experiment = Mock() + log_buffer = launch_nemo_run_workload._ThreadSafeBuffer() + log_buffer.write("Watchdog caught collective operation timeout") + stop_event = threading.Event() + stop_event.set() + failure_detected_event = threading.Event() + + launch_nemo_run_workload._cancel_on_flaky_failure( + experiment, "task-1", log_buffer, stop_event, failure_detected_event, poll_interval=0.01 + ) + + assert not failure_detected_event.is_set() + experiment.cancel.assert_not_called() diff --git a/tests/unit_tests/a2a_overlap/test_cuda_graphed_schedule_chunk_1f1b.py b/tests/unit_tests/a2a_overlap/test_cuda_graphed_schedule_chunk_1f1b.py index 3db52946117..b4351cbbe1e 100644 --- a/tests/unit_tests/a2a_overlap/test_cuda_graphed_schedule_chunk_1f1b.py +++ b/tests/unit_tests/a2a_overlap/test_cuda_graphed_schedule_chunk_1f1b.py @@ -26,19 +26,14 @@ set_global_variables, ) from megatron.training.training import setup_model_and_optimizer +from tests.unit_tests.a2a_overlap.utils import ( + get_valid_flex_dispatcher_backend, + get_valid_token_dispatcher_types, +) from tests.unit_tests.test_utilities import Utils - -def is_deep_ep_available(): - from megatron.core.transformer.moe.fused_a2a import HAVE_DEEP_EP - - return HAVE_DEEP_EP - - -def is_hybrid_ep_available(): - from megatron.core.transformer.moe.fused_a2a import HAVE_HYBRIDEP - - return HAVE_HYBRIDEP +# Transformer Engine 2.17 aborts in the A2A overlap suite with a pybind11 GIL dec_ref failure. +pytestmark = pytest.mark.flaky_in_dev def save(fn, message): @@ -339,22 +334,12 @@ def _run_test_helper( not (HAVE_TE and is_te_min_version("2.10.0")), reason="Partial CUDA graph support requires TransformerEngine version >= 2.10.0", ) - @pytest.mark.parametrize("moe_dispatcher_type", ["alltoall", "deepep"]) + @pytest.mark.parametrize("moe_dispatcher_type", get_valid_token_dispatcher_types()) def test_moe_partial_cudagraph_with_ep_overlap(self, moe_dispatcher_type): extra_kwargs = {"moe_layer_freq": 1} - if moe_dispatcher_type == "deepep": - if not is_deep_ep_available(): - pytest.skip("Deep EP is not available") - extra_kwargs["moe_token_dispatcher_type"] = "flex" - extra_kwargs["moe_flex_dispatcher_backend"] = "deepep" - extra_kwargs["moe_router_dtype"] = "fp32" - elif moe_dispatcher_type == "hybridep": - if not is_hybrid_ep_available(): - pytest.skip("Hybrid EP is not available") - extra_kwargs["moe_token_dispatcher_type"] = "flex" - extra_kwargs["moe_flex_dispatcher_backend"] = "hybridep" - else: - extra_kwargs["moe_token_dispatcher_type"] = moe_dispatcher_type + extra_kwargs["moe_token_dispatcher_type"] = moe_dispatcher_type + if moe_dispatcher_type == "flex": + extra_kwargs["moe_flex_dispatcher_backend"] = get_valid_flex_dispatcher_backend() loss_list_ref = self._run_test_helper(4, "none", None, 3, **extra_kwargs) for cuda_graph_modules in [ diff --git a/tests/unit_tests/a2a_overlap/test_delay_wgrad_compute.py b/tests/unit_tests/a2a_overlap/test_delay_wgrad_compute.py index 01b7768b341..141037bea0a 100644 --- a/tests/unit_tests/a2a_overlap/test_delay_wgrad_compute.py +++ b/tests/unit_tests/a2a_overlap/test_delay_wgrad_compute.py @@ -26,6 +26,9 @@ ) from tests.unit_tests.test_utilities import Utils +# Transformer Engine 2.17 aborts in the A2A overlap suite with a pybind11 GIL dec_ref failure. +pytestmark = pytest.mark.flaky_in_dev + NUM_STEPS = 3 SEQ_LEN = 128 VOCAB_SIZE = 512 diff --git a/tests/unit_tests/a2a_overlap/test_fsdp_1f1b_overlap.py b/tests/unit_tests/a2a_overlap/test_fsdp_1f1b_overlap.py index ec6043f59df..96c2389019a 100644 --- a/tests/unit_tests/a2a_overlap/test_fsdp_1f1b_overlap.py +++ b/tests/unit_tests/a2a_overlap/test_fsdp_1f1b_overlap.py @@ -27,6 +27,9 @@ ) from tests.unit_tests.test_utilities import Utils +# Transformer Engine 2.17 aborts in the A2A overlap suite with a pybind11 GIL dec_ref failure. +pytestmark = pytest.mark.flaky_in_dev + SEQ_LEN = 32 VOCAB_SIZE = 128 NUM_STEPS = 3 diff --git a/tests/unit_tests/a2a_overlap/test_mhc_schedule.py b/tests/unit_tests/a2a_overlap/test_mhc_schedule.py index 4370166f44b..29f99171ccf 100644 --- a/tests/unit_tests/a2a_overlap/test_mhc_schedule.py +++ b/tests/unit_tests/a2a_overlap/test_mhc_schedule.py @@ -88,7 +88,7 @@ class _RecordingLayer: def __init__(self, calls, prefix): self.calls = calls self.config = SimpleNamespace(ep_overlap_early_attn_memory_release=False) - self.attn = _RecordingNode(calls, f"{prefix}.attn") + self.pre_dispatch_computation = _RecordingNode(calls, f"{prefix}.pre_dispatch_computation") self.moe_dispatch = _RecordingNode(calls, f"{prefix}.moe_dispatch") self.mlp = _RecordingNode(calls, f"{prefix}.mlp") self.moe_combine = _RecordingNode(calls, f"{prefix}.moe_combine") @@ -99,7 +99,7 @@ def get_fp8_context(self): return nullcontext() def release_state(self): - self.calls.append(f"{self.attn.name.split('.')[0]}.release_state") + self.calls.append(f"{self.pre_dispatch_computation.name.split('.')[0]}.release_state") class _RecordingChunk: diff --git a/tests/unit_tests/a2a_overlap/test_schedule_chunk_1f1b.py b/tests/unit_tests/a2a_overlap/test_schedule_chunk_1f1b.py index 30fd78c0649..bf8cf45f6be 100644 --- a/tests/unit_tests/a2a_overlap/test_schedule_chunk_1f1b.py +++ b/tests/unit_tests/a2a_overlap/test_schedule_chunk_1f1b.py @@ -23,6 +23,9 @@ ) from tests.unit_tests.test_utilities import Utils +# Transformer Engine 2.17 aborts in the A2A overlap suite with a pybind11 GIL dec_ref failure. +pytestmark = pytest.mark.flaky_in_dev + def build_model(config, use_padding_mask=False): seq_len = 32 diff --git a/tests/unit_tests/a2a_overlap/test_schedule_layer_1f1b.py b/tests/unit_tests/a2a_overlap/test_schedule_layer_1f1b.py index d1bb97ca0cd..fe36e052e03 100644 --- a/tests/unit_tests/a2a_overlap/test_schedule_layer_1f1b.py +++ b/tests/unit_tests/a2a_overlap/test_schedule_layer_1f1b.py @@ -3,6 +3,7 @@ import pytest import torch +import torch.nn.functional as F from megatron.core.fp8_utils import get_fp8_context from megatron.core.models.common.model_chunk_schedule_plan import TransformerLayerSchedulePlan @@ -27,6 +28,32 @@ ) from tests.unit_tests.test_utilities import Utils +# Transformer Engine 2.17 aborts in the A2A overlap suite with a pybind11 GIL dec_ref failure. +pytestmark = pytest.mark.flaky_in_dev + + +def is_nccl_ep_zero_copy_available(): + """Zero-copy needs the newer TE symm-mem APIs (symm_mem_alloc/is_symm_backed), absent in a plain + NCCL-EP build.""" + from megatron.core.transformer.moe.fused_a2a import HAVE_TE_EP + + if not HAVE_TE_EP: + return False + try: + from transformer_engine.pytorch.ep import is_symm_backed, symm_mem_alloc # noqa: F401 + except ImportError: + return False + return True + + +def is_op_fuser_available(): + """The static-shape/zero-copy path runs the TE op-fuser grouped GEMM (needs TE>=2.14 ops).""" + try: + from transformer_engine.pytorch.ops import GroupedLinear, ScaledSwiGLU # noqa: F401 + except ImportError: + return False + return is_te_min_version("2.14.0") + def run_transformer_layer_ref_with_capture(model, input_tensors, iterations): """ @@ -442,6 +469,79 @@ def test_transformer_layer_overlap(self, dispatcher_type, flex_backend, fp8_flag comp_res = compare_captures(capture_ref, capture_a2a_overlap, True) assert comp_res[0], f"[rank {torch.distributed.get_rank()}] {comp_res[1]}" + @pytest.mark.skipif(not is_te_min_version("1.9.0.dev0"), reason="Requires TE >= 1.9.0.dev0") + @pytest.mark.skipif( + not is_nccl_ep_zero_copy_available(), reason="NCCL EP zero-copy TE API is not available" + ) + @pytest.mark.skipif( + not is_op_fuser_available(), reason="op-fuser (static-shape/zero-copy) needs TE>=2.14" + ) + def test_transformer_layer_overlap_zero_copy(self): + """ncclEP zero-copy under 1F1B a2a overlap must match the non-overlap reference. + + Zero-copy stays enabled in both runs, so this isolates the overlap schedule. It also + compares the two ways zero-copy makes the dispatch-backward gradient symm-mem-backed: + the reference gets it from the op-fuser's ``grad_input_buffer``, the overlap run from + ``StageDispatchBwdGrad`` staging into the same buffer (plus the free_input symm guard). + bf16 op-fuser (SwiGLU, tp=1) -- no fp8/Blackwell dependency. + """ + extra_kwargs = {} + apply_flex_backend_kwargs(extra_kwargs, "flex", "ncclep") + extra_kwargs.update( + moe_ncclep_zero_copy=True, + moe_ncclep_static_shape=True, + use_transformer_engine_op_fuser=True, + gated_linear_unit=True, + activation_func=F.silu, + overlap_moe_expert_parallel_comm=True, + ) + config = get_test_config(extra_kwargs=extra_kwargs) + microbatches = 4 + from megatron.core.transformer.moe.fused_a2a import nccl_ep_finalize + from megatron.core.transformer.moe.token_dispatcher import _NCCLEPManager + + try: + with deterministic_mode(): + transformer_layer_spec = get_gpt_decoder_block_spec( + config=config, use_transformer_engine=True + ) + gpt_model = GPTModel( + config=config, + transformer_layer_spec=transformer_layer_spec, + vocab_size=100, + pre_process=True, + post_process=True, + max_sequence_length=300, + ) + params = reset_model(gpt_model) + input_tensors = [build_data() for _ in range(microbatches)] + + # The reference runs the layer directly instead of through the 1F1B schedule, so it + # must declare overlap=False: that is what makes get_expert_zero_copy_buffers hand + # the op-fuser the symm grad_input_buffer for the fc1 dgrad. Under overlap=True the + # buffer is withheld (the schedule detaches the dispatch output, so autograd would + # discard it) and StageDispatchBwdGrad supplies the symm gradient instead. + config.overlap_moe_expert_parallel_comm = False + capture_ref = run_transformer_layer_ref_with_capture( + gpt_model, input_tensors, microbatches + ) + config.overlap_moe_expert_parallel_comm = True + + reset_model(gpt_model, params) + capture_a2a_overlap = run_transformer_layer_a2a_overlap_with_capture( + gpt_model, input_tensors, microbatches + ) + comp_res = compare_captures(capture_ref, capture_a2a_overlap, True) + assert comp_res[0], f"[rank {torch.distributed.get_rank()}] {comp_res[1]}" + finally: + # zero-copy sets process-global ncclEP state (ep bootstrap mode + shared symm + # classvars). Reset in a finally: on failure the leaked classvars would otherwise make + # every later ncclEP test in this process fail too, hiding the real error. + nccl_ep_finalize() + _NCCLEPManager._zc_fwd_token_buf = None + _NCCLEPManager._zc_bwd_token_buf = None + _NCCLEPManager._zc_recv_topk_weights_buf = None + @pytest.mark.skipif(not is_te_min_version("1.9.0.dev0"), reason="Requires TE >= 1.9.0.dev0") @pytest.mark.parametrize("dispatcher_type,flex_backend", get_valid_dispatcher_configs()) @pytest.mark.parametrize("fp8_flag", get_valid_fp8_flags()) @@ -450,7 +550,12 @@ def test_mtp_layer_overlap(self, dispatcher_type, flex_backend, fp8_flag): Verifies all-to-all overlap optimization in MTP layer produces the same results as the reference implementation. """ - extra_kwargs = {"mtp_num_layers": 1, "mtp_loss_scaling_factor": 1.1} + qk_layernorm = True + extra_kwargs = { + "mtp_num_layers": 1, + "mtp_loss_scaling_factor": 1.1, + "qk_layernorm": qk_layernorm, + } apply_flex_backend_kwargs(extra_kwargs, dispatcher_type, flex_backend) if fp8_flag is not None: extra_kwargs["fp8_recipe"] = fp8_flag[1] @@ -463,7 +568,7 @@ def test_mtp_layer_overlap(self, dispatcher_type, flex_backend, fp8_flag): transformer_layer_spec = get_gpt_layer_with_transformer_engine_spec( num_experts=16, moe_grouped_gemm=True, - qk_layernorm=True, + qk_layernorm=qk_layernorm, multi_latent_attention=True, ) mtp_block_spec = get_gpt_mtp_block_spec(config, transformer_layer_spec, True) diff --git a/tests/unit_tests/conftest.py b/tests/unit_tests/conftest.py index ef3d87c7c6d..4620c6a57fc 100644 --- a/tests/unit_tests/conftest.py +++ b/tests/unit_tests/conftest.py @@ -1,7 +1,6 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import os -from datetime import timedelta from pathlib import Path import pytest @@ -15,6 +14,22 @@ from tests.unit_tests.test_utilities import Utils +def pytest_configure(config): + """Set NCCL defaults for the unit-test suite. + + These previously lived as ``export``s in ``tests/unit_tests/run_ci_test.sh``. + They reduce NCCL memory usage / SM contention and were originally added to + fix NCCL hangs observed for FSDP v1 (among other MCore algorithms). Setting + them here — at session start, before any test initializes NCCL communicators + — keeps that default while moving the test-bucket configuration out of the + CI launch script and into pytest. Individual buckets that want + production-like NCCL settings (e.g. MFSDP v2) can pop these in their own + conftest before initializing their process group. + """ + os.environ.setdefault("NCCL_MAX_NCHANNELS", "1") + os.environ.setdefault("NCCL_NVLS_ENABLE", "0") + + def pytest_addoption(parser): """ Additional command-line arguments passed to pytest. @@ -44,7 +59,10 @@ def cleanup(): yield if torch.distributed.is_initialized(): try: - torch.distributed.barrier() + if torch.cuda.is_available(): + torch.distributed.barrier(device_ids=[torch.cuda.current_device()]) + else: + torch.distributed.barrier() except Exception: return torch.distributed.destroy_process_group() diff --git a/tests/unit_tests/dist_checkpointing/models/test_gpt_hybrid_interop.py b/tests/unit_tests/dist_checkpointing/models/test_gpt_hybrid_interop.py new file mode 100644 index 00000000000..dae973f4a7b --- /dev/null +++ b/tests/unit_tests/dist_checkpointing/models/test_gpt_hybrid_interop.py @@ -0,0 +1,827 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Tests for loading GPT checkpoints into HybridModel runs. + +Covers the load-time sharded state dict retargeting in +``megatron.core.dist_checkpointing.gpt_checkpoint_interop``: + +* pure layer-map derivation and validation (no GPU state), +* pure key retargeting on synthetic sharded state dicts, +* end-to-end: save a GPTModel dist checkpoint under one (TP, PP, EP, ETP) + layout, load it into a HybridModel under another layout, and verify + attention/MLP weights round-trip bit-for-bit while layers without a GPT + counterpart keep their fresh initialization. +""" + +from functools import partial +from unittest import mock + +import pytest +import torch + +from megatron.core import parallel_state as ps +from megatron.core.dist_checkpointing import load, load_plain_tensors, save +from megatron.core.dist_checkpointing.dict_utils import diff +from megatron.core.dist_checkpointing.gpt_checkpoint_interop import ( + gpt_compatible_layer_maps, + retarget_fsdp_state_dict_to_gpt_checkpoint, + retarget_sharded_state_dict_to_gpt_checkpoint, +) +from megatron.core.dist_checkpointing.mapping import ( + LocalNonpersistentObject, + ShardedObject, + ShardedTensor, + ShardedTensorFactory, +) +from megatron.core.dist_checkpointing.validation import StrictHandling +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_decoder_block_spec, + get_gpt_layer_with_transformer_engine_spec, +) +from megatron.core.models.gpt.gpt_model import GPTModel +from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec +from megatron.core.models.hybrid.hybrid_model import HybridModel +from megatron.core.num_microbatches_calculator import ( + destroy_num_microbatches_calculator, + init_num_microbatches_calculator, +) +from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.training.arguments import parse_args +from megatron.training.checkpointing import load_checkpoint, save_checkpoint +from tests.unit_tests.dist_checkpointing import TempNamedDir +from tests.unit_tests.dist_checkpointing.utils import ( + init_checkpointing_mock_args, + setup_model_and_optimizer, +) +from tests.unit_tests.test_utilities import Utils + + +class TestGPTCompatLayerMaps: + def test_pairs_positions_in_pattern_order(self): + maps = gpt_compatible_layer_maps('M*-M*-') + assert maps.attention_to_gpt == {1: 0, 4: 1} + assert maps.mlp_to_gpt == {2: 0, 5: 1} + assert maps.fresh_init == frozenset({0, 3}) + assert maps.num_gpt_layers == 2 + + def test_pipeline_separators_are_ignored(self): + assert gpt_compatible_layer_maps('M*-|M*-') == gpt_compatible_layer_maps('M*-M*-') + + def test_moe_positions_pair_like_dense_ones(self): + maps = gpt_compatible_layer_maps('M*EM*E') + assert maps.attention_to_gpt == {1: 0, 4: 1} + assert maps.mlp_to_gpt == {2: 0, 5: 1} + assert maps.num_gpt_layers == 2 + + def test_attention_only_positions_can_precede_all_mlps(self): + # Pairing is positional per type, not adjacency-based. + maps = gpt_compatible_layer_maps('**--') + assert maps.attention_to_gpt == {0: 0, 1: 1} + assert maps.mlp_to_gpt == {2: 0, 3: 1} + assert maps.fresh_init == frozenset() + + def test_rejects_empty_pattern(self): + with pytest.raises(ValueError, match='empty'): + gpt_compatible_layer_maps(None) + + def test_rejects_mtp_pattern(self): + with pytest.raises(ValueError, match='MTP'): + gpt_compatible_layer_maps('M*-M*-/MM/MM') + + def test_rejects_untranslatable_layer_types(self): + with pytest.raises(ValueError, match='cannot be translated'): + gpt_compatible_layer_maps('M*-G*-') + # 'D' cannot be combined with '*' at all, so use a pattern the + # production parser accepts and let the interop validation reject it. + with pytest.raises(ValueError, match='cannot be translated'): + gpt_compatible_layer_maps('MD-') + + def test_rejects_mixed_dense_and_moe(self): + with pytest.raises(ValueError, match="one of '-' or 'E'"): + gpt_compatible_layer_maps('M*-M*E') + + def test_rejects_unbalanced_attention_and_mlp(self): + with pytest.raises(ValueError, match='equal, nonzero'): + gpt_compatible_layer_maps('M**-') + with pytest.raises(ValueError, match='equal, nonzero'): + gpt_compatible_layer_maps('MMMM') + + +def _sharded_tensor(key): + return ShardedTensor.from_rank_offsets(key, torch.ones(4)) + + +class TestRetargetShardedStateDict: + def test_keys_point_at_gpt_layout_and_fresh_layers_stay_local(self): + # GPT checkpoints use the homogeneous layer format: numberless keys + # with the layer index as the leading sharding axis. + maps = gpt_compatible_layer_maps('M*-') + mixer_weight = torch.full((4,), 7.0) + sharded_sd = { + 'attn': _sharded_tensor('decoder.layers.1.self_attention.linear_qkv.weight'), + 'mlp': _sharded_tensor('decoder.layers.2.mlp.linear_fc1.layer_norm_weight'), + 'mixer': ShardedTensor.from_rank_offsets( + 'decoder.layers.0.mixer.in_proj.weight', mixer_weight + ), + 'final_norm': _sharded_tensor('decoder.final_norm.weight'), + 'embedding': _sharded_tensor('embedding.word_embeddings.weight'), + 'output_extra_state': ShardedObject('output_layer._extra_state', None, (1,), (0,)), + 'nested': { + 'proj': _sharded_tensor('decoder.layers.1.self_attention.linear_proj.weight') + }, + } + + retarget_sharded_state_dict_to_gpt_checkpoint(sharded_sd, maps) + + attn = sharded_sd['attn'] + assert attn.key == 'decoder.layers.self_attention.linear_qkv.weight' + assert attn.prepend_axis_num == 1 + assert attn.global_shape == (1, 4) + assert attn.global_offset == (0, 0) + assert attn.axis_fragmentations == (1, 1) + assert sharded_sd['mlp'].key == 'decoder.layers.mlp.linear_fc1.layer_norm_weight' + assert sharded_sd['mlp'].global_offset == (0, 0) + assert ( + sharded_sd['nested']['proj'].key == 'decoder.layers.self_attention.linear_proj.weight' + ) + assert sharded_sd['final_norm'].key == 'decoder.final_layernorm.weight' + assert sharded_sd['final_norm'].prepend_axis_num == 0 + assert sharded_sd['embedding'].key == 'embedding.word_embeddings.weight' + assert isinstance(sharded_sd['output_extra_state'], LocalNonpersistentObject) + assert sharded_sd['output_extra_state'].unwrap() is None + assert isinstance(sharded_sd['mixer'], LocalNonpersistentObject) + assert sharded_sd['mixer'].unwrap() is mixer_weight + + def test_extra_state_and_factory_entries_follow_the_layer_axis(self): + maps = gpt_compatible_layer_maps('M*-M*-') # 2 GPT layers + extra_state = ShardedObject( + 'decoder.layers.5.mlp.linear_fc2._extra_state', None, (1,), (0,) + ) + + def build_fn(key, data, replica_id, flattened_range): + return { + 'chunk': ShardedTensor.from_rank_offsets( + f'{key}_chunk', data, replica_id=replica_id + ) + } + + factory = ShardedTensorFactory( + 'decoder.layers.2.mlp.linear_fc1.weight', + torch.ones(4), + build_fn, + lambda sd: sd['chunk'], + ) + sharded_sd = {'extra_state': extra_state, 'factory': factory} + + retarget_sharded_state_dict_to_gpt_checkpoint(sharded_sd, maps) + + # Hybrid layer 5 is the 2nd MLP position -> GPT layer 1. + assert extra_state.key == 'decoder.layers.mlp.linear_fc2._extra_state' + assert extra_state.global_shape == (2,) + assert extra_state.global_offset == (1,) + # Hybrid layer 2 is the 1st MLP position -> GPT layer 0; sub-tensors + # built by the factory inherit the layer axis. + assert factory.key == 'decoder.layers.mlp.linear_fc1.weight' + built = factory.build() + assert built['chunk'].key == 'decoder.layers.mlp.linear_fc1.weight_chunk' + assert built['chunk'].prepend_axis_num == 1 + assert built['chunk'].global_shape == (2, 4) + assert built['chunk'].global_offset == (0, 0) + + def test_layer_outside_pattern_raises(self): + maps = gpt_compatible_layer_maps('M*-') + sharded_sd = {'bad': _sharded_tensor('decoder.layers.7.self_attention.linear_qkv.weight')} + with pytest.raises(ValueError, match='not part of the hybrid layer pattern'): + retarget_sharded_state_dict_to_gpt_checkpoint(sharded_sd, maps) + + def test_optimizer_state_entries_retarget_like_the_model(self): + # The distributed optimizer's model-space sharded state dict embeds the + # model key under ``optimizer.state..`` and mirrors the + # model param's sharding, so the same retargeting must point the moments + # and fp32 master params at the GPT checkpoint and keep fresh-layer + # optimizer state local. + maps = gpt_compatible_layer_maps('M*-') + fresh_exp_avg = torch.full((4,), 3.0) + optim_sd = { + 'param_state': { + # attention position (hybrid layer 1 -> GPT layer 0) + 0: { + 'exp_avg': _sharded_tensor( + 'optimizer.state.exp_avg.decoder.layers.1.self_attention.linear_qkv.weight' + ), + 'fp32_param': _sharded_tensor( + 'optimizer.state.fp32_param.decoder.layers.1.self_attention.linear_qkv.weight' + ), + }, + # MLP position (hybrid layer 2 -> GPT layer 0) + 1: { + 'exp_avg_sq': _sharded_tensor( + 'optimizer.state.exp_avg_sq.decoder.layers.2.mlp.linear_fc1.weight' + ) + }, + # fresh Mamba position (hybrid layer 0) -> stays local + 2: { + 'exp_avg': ShardedTensor.from_rank_offsets( + 'optimizer.state.exp_avg.decoder.layers.0.mixer.in_proj.weight', + fresh_exp_avg, + ) + }, + }, + 'param_state_sharding_type': 'fully_sharded_model_space', + } + + retarget_sharded_state_dict_to_gpt_checkpoint(optim_sd, maps) + + attn = optim_sd['param_state'][0]['exp_avg'] + assert attn.key == 'optimizer.state.exp_avg.decoder.layers.self_attention.linear_qkv.weight' + assert attn.prepend_axis_num == 1 + assert attn.global_shape == (1, 4) + assert attn.global_offset == (0, 0) + master = optim_sd['param_state'][0]['fp32_param'] + assert ( + master.key + == 'optimizer.state.fp32_param.decoder.layers.self_attention.linear_qkv.weight' + ) + assert optim_sd['param_state'][1]['exp_avg_sq'].key == ( + 'optimizer.state.exp_avg_sq.decoder.layers.mlp.linear_fc1.weight' + ) + fresh = optim_sd['param_state'][2]['exp_avg'] + assert isinstance(fresh, LocalNonpersistentObject) + assert fresh.unwrap() is fresh_exp_avg + # Non-sharded bookkeeping is passed through untouched. + assert optim_sd['param_state_sharding_type'] == 'fully_sharded_model_space' + + def test_fsdp_model_and_optimizer_keys_retarget_recursively(self): + maps = gpt_compatible_layer_maps('M*-') + attn = torch.ones(2) + mlp_moment = torch.ones(2) + fresh = torch.ones(2) + state_dict = { + 'model': { + 'decoder.layers.1.self_attention.linear_qkv.weight': attn, + 'decoder.layers.0.mixer.in_proj.weight': fresh, + 'decoder.final_norm.weight': torch.ones(2), + 'output_layer._extra_state': None, + }, + 'optimizer': { + 'state': { + 'decoder.layers.2.mlp.linear_fc1.weight': {'exp_avg': mlp_moment}, + 'decoder.layers.0.mixer.in_proj.weight': {'exp_avg': fresh}, + }, + 'param_to_group_meta': { + 'decoder.layers.2.mlp.linear_fc1.weight': {'lr_mult': 1.0}, + 'decoder.layers.0.mixer.in_proj.weight': {'lr_mult': 1.0}, + }, + }, + } + + translated = retarget_fsdp_state_dict_to_gpt_checkpoint(state_dict, maps) + + assert set(translated['model']) == { + 'decoder.layers.0.self_attention.linear_qkv.weight', + 'decoder.final_layernorm.weight', + } + assert translated['model']['decoder.layers.0.self_attention.linear_qkv.weight'] is attn + assert set(translated['optimizer']['state']) == {'decoder.layers.0.mlp.linear_fc1.weight'} + assert ( + translated['optimizer']['state']['decoder.layers.0.mlp.linear_fc1.weight']['exp_avg'] + is mlp_moment + ) + assert set(translated['optimizer']['param_to_group_meta']) == { + 'decoder.layers.0.mlp.linear_fc1.weight' + } + + wrapped = {'module.module.module.decoder.layers.1.self_attention.linear_qkv.weight': attn} + translated = retarget_fsdp_state_dict_to_gpt_checkpoint( + wrapped, + maps, + ('optimizer.state.module.module.decoder.layers.0.' 'self_attention.linear_qkv.weight',), + checkpoint_prefix='optimizer.state', + ) + assert set(translated) == { + 'module.module.decoder.layers.0.self_attention.linear_qkv.weight' + } + + +def _base_config_kwargs(parallel, moe, glu): + tp, pp, ep, etp = parallel + config_kwargs = dict( + num_attention_heads=8, + # for Mamba: expand=2, headdim=64 -> nheads=8 (divisible by ngroups=8) + hidden_size=256, + use_cpu_initialization=True, + pipeline_dtype=torch.bfloat16, + tensor_model_parallel_size=tp, + pipeline_model_parallel_size=pp, + sequence_parallel=(tp > 1 and ep > 1), + # gated MLPs exercise the swiglu ShardedTensorFactory path + gated_linear_unit=glu, + add_bias_linear=not glu, + ) + if moe: + config_kwargs.update( + num_moe_experts=8, + moe_grouped_gemm=True, # the hybrid moe spec is built with grouped GEMM experts + add_bias_linear=False, + moe_router_topk=2, + expert_model_parallel_size=ep, + expert_tensor_parallel_size=etp, + ) + return config_kwargs + + +def initialize_gpt_model(seed, num_gpt_layers, parallel, moe, glu=False): + torch.manual_seed(seed) + model_parallel_cuda_manual_seed(seed) + + config = TransformerConfig(num_layers=num_gpt_layers, **_base_config_kwargs(parallel, moe, glu)) + if moe: + layer_spec = get_gpt_decoder_block_spec(config, use_transformer_engine=True) + else: + layer_spec = get_gpt_layer_with_transformer_engine_spec() + model = GPTModel( + config=config, + transformer_layer_spec=layer_spec, + vocab_size=128, + max_sequence_length=4, + pre_process=ps.is_pipeline_first_stage(), + post_process=ps.is_pipeline_last_stage(), + position_embedding_type='rope', + share_embeddings_and_output_weights=True, + ) + with torch.no_grad(): + for param in model.parameters(): + param.random_() + return model + + +def initialize_hybrid_model(seed, pattern, parallel, moe, glu=False): + torch.manual_seed(seed) + model_parallel_cuda_manual_seed(seed) + + num_layers = len(pattern.replace('|', '')) + config = TransformerConfig(num_layers=num_layers, **_base_config_kwargs(parallel, moe, glu)) + return HybridModel( + config=config, + hybrid_stack_spec=hybrid_stack_spec, + vocab_size=128, + max_sequence_length=4, + hybrid_layer_pattern=pattern, + pre_process=ps.is_pipeline_first_stage(), + post_process=ps.is_pipeline_last_stage(), + position_embedding_type='rope', + share_embeddings_and_output_weights=True, + ) + + +def _snapshot_fresh_layers(hybrid_model, layer_maps): + """Clone all tensors of layers that must keep their fresh initialization.""" + snapshot = {} + for layer in hybrid_model.decoder.layers: + global_idx = layer.layer_number - 1 + if global_idx in layer_maps.fresh_init: + snapshot[global_idx] = { + name: tensor.detach().clone() + for name, tensor in layer.state_dict().items() + if isinstance(tensor, torch.Tensor) + } + return snapshot + + +def _assert_fresh_layers_untouched(hybrid_model, layer_maps, snapshot): + for layer in hybrid_model.decoder.layers: + global_idx = layer.layer_number - 1 + if global_idx not in layer_maps.fresh_init: + continue + for name, tensor in layer.state_dict().items(): + if not isinstance(tensor, torch.Tensor): + continue + assert torch.equal( + tensor, snapshot[global_idx][name] + ), f'fresh layer {global_idx} tensor {name} was overwritten by the GPT load' + + +def _drop_extra_state(plain_state_dict): + return {k: v for k, v in plain_state_dict.items() if '_extra_state' not in k} + + +class TestGPTToHybridLoad: + def teardown_method(self, method): + Utils.destroy_model_parallel() + + @pytest.mark.internal + @pytest.mark.parametrize( + ('src_parallel', 'dest_parallel', 'pattern', 'moe', 'glu'), + [ + # (tp, pp, ep, etp) of the GPT save -> of the hybrid load. + # Dense: TP/PP resharding. + ((1, 1, 1, 1), (1, 1, 1, 1), 'M*-M*-M*-M*-', False, False), + ((2, 1, 1, 1), (1, 2, 1, 1), 'M*-M*-M*-M*-', False, False), + ((1, 2, 1, 1), (2, 1, 1, 1), 'M*-M*-M*-M*-', False, False), + ((2, 2, 1, 1), (4, 1, 1, 1), 'M*-M*-M*-M*-', False, False), + ((4, 1, 1, 1), (2, 4, 1, 1), 'M*-M*-M*-M*-', False, False), + # Gated MLP exercises the swiglu factory path. + ((2, 1, 1, 1), (1, 2, 1, 1), 'M*-M*-M*-M*-', False, True), + # Pipeline stage boundaries given explicitly with '|'. + ((1, 2, 1, 1), (1, 2, 1, 1), 'M*-M*-|M*-M*-', False, False), + # MoE: EP/ETP resharding (ETP defaults to TP when 1). + ((1, 1, 1, 1), (1, 1, 4, 1), 'M*EM*E', True, False), + ((1, 1, 4, 1), (2, 1, 1, 2), 'M*EM*E', True, False), + ((2, 1, 2, 2), (1, 1, 8, 1), 'M*EM*E', True, False), + ((1, 1, 2, 1), (4, 1, 2, 4), 'M*EM*E', True, False), + # MoE with PP as well. + ((2, 1, 2, 2), (1, 2, 2, 1), 'M*EM*EM*EM*E', True, False), + ], + ) + def test_gpt_checkpoint_loads_into_hybrid_across_parallel_layouts( + self, tmp_path_dist_ckpt, src_parallel, dest_parallel, pattern, moe, glu + ): + layer_maps = gpt_compatible_layer_maps(pattern) + src_tp, src_pp, src_ep, src_etp = src_parallel + dest_tp, dest_pp, dest_ep, dest_etp = dest_parallel + + Utils.initialize_model_parallel( + src_tp, src_pp, expert_model_parallel_size=src_ep, expert_tensor_parallel_size=src_etp + ) + with ( + TempNamedDir(tmp_path_dist_ckpt / 'gpt_hybrid_interop_gpt_src') as ckpt_dir_gpt, + TempNamedDir(tmp_path_dist_ckpt / 'gpt_hybrid_interop_roundtrip') as ckpt_dir_back, + ): + # Save a GPT checkpoint under the source parallel layout. + gpt_model = initialize_gpt_model(1, layer_maps.num_gpt_layers, src_parallel, moe, glu) + save(gpt_model.sharded_state_dict(), ckpt_dir_gpt) + Utils.destroy_model_parallel() + + # Load it into a hybrid model under the destination layout by + # retargeting the hybrid sharded state dict, exactly as + # load_checkpoint does for GPT checkpoints. + Utils.initialize_model_parallel( + dest_tp, + dest_pp, + expert_model_parallel_size=dest_ep, + expert_tensor_parallel_size=dest_etp, + ) + hybrid_model = initialize_hybrid_model(2, pattern, dest_parallel, moe, glu) + fresh_snapshot = _snapshot_fresh_layers(hybrid_model, layer_maps) + + sharded_sd = hybrid_model.sharded_state_dict() + retarget_sharded_state_dict_to_gpt_checkpoint(sharded_sd, layer_maps) + state_dict, missing_keys, unexpected_keys = load( + sharded_sd, ckpt_dir_gpt, strict=StrictHandling.RETURN_ALL + ) + # Any mismatch beyond TE extra states means the retargeting missed keys. + assert all('_extra_state' in k for k in missing_keys), missing_keys + assert all('_extra_state' in k for k in unexpected_keys), unexpected_keys + hybrid_model.load_state_dict(state_dict) + + _assert_fresh_layers_untouched(hybrid_model, layer_maps, fresh_snapshot) + + # Save the hybrid model back under GPT keys (fresh layers stay + # local and are skipped) and compare both checkpoints tensorwise. + sharded_sd_back = hybrid_model.sharded_state_dict() + retarget_sharded_state_dict_to_gpt_checkpoint(sharded_sd_back, layer_maps) + save(sharded_sd_back, ckpt_dir_back) + Utils.destroy_model_parallel() + + Utils.initialize_model_parallel(1, 1) + plain_gpt = _drop_extra_state(load_plain_tensors(ckpt_dir_gpt)) + plain_back = _drop_extra_state(load_plain_tensors(ckpt_dir_back)) + only_gpt, only_back, mismatch = diff(plain_gpt, plain_back) + assert not only_back, f'roundtrip produced keys missing from the GPT ckpt: {only_back}' + assert not only_gpt, f'GPT ckpt keys not covered by the hybrid load: {only_gpt}' + assert not mismatch, f'weights changed by the GPT->hybrid->GPT roundtrip: {mismatch}' + + +# --------------------------------------------------------------------------- +# End-to-end optimizer loading through save_checkpoint / load_checkpoint. +# --------------------------------------------------------------------------- + +_OPT_HIDDEN = 256 +_OPT_HEADS = 8 + + +def _opt_provider_config(num_layers, moe=False, **config_kwargs): + # get_model passes these through; they are not TransformerConfig fields. + for extra in ('pg_collection', 'config', 'vp_stage'): + config_kwargs.pop(extra, None) + config_kwargs.update( + num_layers=num_layers, + hidden_size=_OPT_HIDDEN, + num_attention_heads=_OPT_HEADS, + use_cpu_initialization=True, + add_bias_linear=not moe, + gated_linear_unit=False, + ) + if moe: + config_kwargs.update( + num_moe_experts=8, + moe_grouped_gemm=True, + moe_router_topk=2, + sequence_parallel=( + config_kwargs['tensor_model_parallel_size'] > 1 + and config_kwargs['expert_model_parallel_size'] > 1 + ), + ) + return TransformerConfig(**config_kwargs) + + +def gpt_provider_for_opt( + pre_process=True, post_process=True, *, seed=0, num_gpt_layers, moe=False, **kw +): + torch.manual_seed(seed) + model_parallel_cuda_manual_seed(seed) + config = _opt_provider_config(num_gpt_layers, moe=moe, **kw) + layer_spec = ( + get_gpt_decoder_block_spec(config, use_transformer_engine=True) + if moe + else get_gpt_layer_with_transformer_engine_spec() + ) + return GPTModel( + config=config, + transformer_layer_spec=layer_spec, + vocab_size=128, + max_sequence_length=4, + pre_process=pre_process, + post_process=post_process, + position_embedding_type='rope', + share_embeddings_and_output_weights=True, + ) + + +def hybrid_provider_for_opt( + pre_process=True, post_process=True, *, seed=0, pattern, moe=False, **kw +): + torch.manual_seed(seed) + model_parallel_cuda_manual_seed(seed) + num_layers = len(pattern.replace('|', '')) + return HybridModel( + config=_opt_provider_config(num_layers, moe=moe, **kw), + hybrid_stack_spec=hybrid_stack_spec, + vocab_size=128, + max_sequence_length=4, + hybrid_layer_pattern=pattern, + pre_process=pre_process, + post_process=post_process, + position_embedding_type='rope', + share_embeddings_and_output_weights=True, + ) + + +def _inner_optimizers(optimizer): + if hasattr(optimizer, 'chained_optimizers'): + return [o for opt in optimizer.chained_optimizers for o in _inner_optimizers(opt)] + inner = getattr(optimizer, 'optimizer', None) + return [inner] if inner is not None else [] + + +def _optimizer_moment_fingerprint(optimizer): + """Sum of norms of every floating-point optimizer-state tensor (per rank).""" + total = 0.0 + for inner in _inner_optimizers(optimizer): + for state in inner.state.values(): + for value in state.values(): + if torch.is_tensor(value) and value.is_floating_point(): + total += value.detach().double().norm().item() + return total + + +def _seed_optimizer_moments(optimizer, seed): + """Populate Adam moments even when DistOpt is nested in a chained optimizer.""" + torch.manual_seed(seed) + for inner in _inner_optimizers(optimizer): + for group in inner.param_groups: + for param in group['params']: + state = inner.state[param] + state['exp_avg'] = torch.rand_like(param) + state['exp_avg_sq'] = torch.rand_like(param) + + +def _set_checkpoint_parallel_args(args, parallel, moe): + tp, pp, cp, ep, etp = parallel + args.world_size = torch.distributed.get_world_size() + args.data_parallel_size = ps.get_data_parallel_world_size() + args.tensor_model_parallel_size = tp + args.pipeline_model_parallel_size = pp + args.context_parallel_size = cp + args.expert_model_parallel_size = ep + args.expert_tensor_parallel_size = etp + args.num_experts = 8 if moe else None + + +def _model_parameter_snapshot(model): + return {name: param.detach().clone() for name, param in model.named_parameters()} + + +def _unwrap_interop_model(model): + """Reach GPTModel/HybridModel through Float16 and Megatron-FSDP wrappers.""" + while not hasattr(model, 'decoder') and hasattr(model, 'module'): + model = model.module + return model + + +def _configure_checkpoint_args(args, ckpt_dir, parallel, moe, use_megatron_fsdp): + init_checkpointing_mock_args(args, ckpt_dir, fully_parallel=not use_megatron_fsdp) + _set_checkpoint_parallel_args(args, parallel, moe) + args.use_distributed_optimizer = True + args.use_megatron_fsdp = use_megatron_fsdp + args.data_parallel_sharding_strategy = 'optim_grads_params' if use_megatron_fsdp else 'no_shard' + args.ckpt_format = 'fsdp_dtensor' if use_megatron_fsdp else 'torch_dist' + args.dist_ckpt_optim_fully_reshardable = not use_megatron_fsdp + args.hidden_size = _OPT_HIDDEN + args.num_attention_heads = _OPT_HEADS + + +def _run_gpt_to_hybrid_optimizer_load( + tmp_path_dist_ckpt, + src_parallel, + dest_parallel, + pattern, + moe, + *, + finetune=True, + use_megatron_fsdp=False, +): + layer_maps = gpt_compatible_layer_maps(pattern) + num_gpt_layers = layer_maps.num_gpt_layers + + src_tp, src_pp, src_cp, src_ep, src_etp = src_parallel + dest_tp, dest_pp, dest_cp, dest_ep, dest_etp = dest_parallel + + Utils.initialize_model_parallel( + src_tp, + src_pp, + context_parallel_size=src_cp, + expert_model_parallel_size=src_ep, + expert_tensor_parallel_size=src_etp, + ) + with TempNamedDir(tmp_path_dist_ckpt / 'gpt_hybrid_opt_interop') as ckpt_dir: + mock_args = parse_args(ignore_unknown_args=True) + with mock.patch('megatron.training.checkpointing.get_args', new=lambda: mock_args): + # Build a GPT model + distributed optimizer whose Adam moments are + # seeded to random values, then save a full checkpoint. + gpt_model, gpt_optimizer = setup_model_and_optimizer( + seed=2, + tp=src_tp, + pp=src_pp, + cp=src_cp, + ep=src_ep, + etp=src_etp, + use_megatron_fsdp=use_megatron_fsdp, + initialize_fn=partial(gpt_provider_for_opt, num_gpt_layers=num_gpt_layers, moe=moe), + ) + _seed_optimizer_moments(gpt_optimizer, seed=3) + _configure_checkpoint_args(mock_args, ckpt_dir, src_parallel, moe, use_megatron_fsdp) + mock_args.num_layers = num_gpt_layers + save_checkpoint(10, gpt_model, gpt_optimizer, None, 0) + Utils.destroy_model_parallel() + + # Build a hybrid model + optimizer (independently seeded moments) and + # load the GPT checkpoint, translating model and optimizer state. + Utils.initialize_model_parallel( + dest_tp, + dest_pp, + context_parallel_size=dest_cp, + expert_model_parallel_size=dest_ep, + expert_tensor_parallel_size=dest_etp, + ) + hybrid_model, hybrid_optimizer = setup_model_and_optimizer( + seed=4, + tp=dest_tp, + pp=dest_pp, + cp=dest_cp, + ep=dest_ep, + etp=dest_etp, + use_megatron_fsdp=use_megatron_fsdp, + initialize_fn=partial(hybrid_provider_for_opt, pattern=pattern, moe=moe), + ) + _seed_optimizer_moments(hybrid_optimizer, seed=5) + hybrid_module = _unwrap_interop_model(hybrid_model[0]) + fresh_snapshot = _snapshot_fresh_layers(hybrid_module, layer_maps) + model_before = _model_parameter_snapshot(hybrid_module) + moments_before = _optimizer_moment_fingerprint(hybrid_optimizer) + + _configure_checkpoint_args(mock_args, ckpt_dir, dest_parallel, moe, use_megatron_fsdp) + mock_args.finetune = finetune + mock_args.hybrid_layer_pattern = pattern + mock_args.num_layers = len(pattern.replace('|', '')) + + if not finetune: + data_parallel_size = ps.get_data_parallel_world_size() + init_num_microbatches_calculator( + rank=torch.distributed.get_rank(), + global_batch_size=data_parallel_size, + micro_batch_size=1, + data_parallel_size=data_parallel_size, + ) + try: + iteration, _ = load_checkpoint(hybrid_model, hybrid_optimizer, None) + finally: + if not finetune: + destroy_num_microbatches_calculator() + + # GPT-to-Hybrid translation does not select checkpoint semantics: + # --finetune restarts iteration, while a regular load resumes it. + assert iteration == (0 if finetune else 10) + assert any( + not torch.equal(param, model_before[name]) + for name, param in hybrid_module.named_parameters() + ), 'model parameters do not appear to have been loaded' + # The optimizer state was actually loaded (GPT-sourced moments + # overwrite the freshly seeded ones). + moments_after = _optimizer_moment_fingerprint(hybrid_optimizer) + assert ( + abs(moments_after - moments_before) > 1e-6 + ), 'optimizer state does not appear to have been loaded' + # Layers without a GPT counterpart keep their fresh weights. + _assert_fresh_layers_untouched(hybrid_module, layer_maps, fresh_snapshot) + + +class TestGPTToHybridOptimizerLoad: + def teardown_method(self, method): + Utils.destroy_model_parallel() + + @pytest.mark.internal + @pytest.mark.parametrize( + ('src_parallel', 'dest_parallel', 'pattern', 'moe', 'finetune'), + [ + pytest.param( + (1, 1, 1, 1, 1), + (1, 1, 1, 1, 1), + 'M*-M*-', + False, + False, + id='resume-without-finetune', + ), + pytest.param((1, 1, 1, 1, 1), (1, 1, 1, 1, 1), '*-*-', False, True, id='dp8-dense'), + pytest.param((1, 1, 1, 1, 1), (1, 1, 1, 1, 1), 'M*-M*-', False, True, id='dp8-hybrid'), + pytest.param( + (1, 1, 1, 1, 1), (2, 1, 1, 1, 1), 'M*-M*-', False, True, id='dp8-to-tp2-dp4' + ), + pytest.param( + (2, 1, 1, 1, 1), (1, 2, 1, 1, 1), 'M*-M*-', False, True, id='tp2-dp4-to-pp2-dp4' + ), + pytest.param( + (1, 1, 2, 1, 1), (2, 1, 1, 1, 1), 'M*-M*-', False, True, id='cp2-dp4-to-tp2-dp4' + ), + pytest.param( + (1, 1, 2, 4, 1), (2, 1, 1, 2, 2), 'M*EM*E', True, True, id='cp2-ep4-to-tp2-ep2-etp2' + ), + pytest.param( + (2, 1, 1, 2, 2), (1, 2, 1, 4, 1), 'M*EM*E', True, True, id='tp2-ep2-etp2-to-pp2-ep4' + ), + ], + ) + def test_gpt_optimizer_state_loads_into_hybrid( + self, tmp_path_dist_ckpt, src_parallel, dest_parallel, pattern, moe, finetune + ): + _run_gpt_to_hybrid_optimizer_load( + tmp_path_dist_ckpt, src_parallel, dest_parallel, pattern, moe, finetune=finetune + ) + + +class TestGPTToHybridFSDPLoad: + def teardown_method(self, method): + Utils.destroy_model_parallel() + + @pytest.mark.internal + @pytest.mark.parametrize( + ('src_parallel', 'dest_parallel', 'pattern', 'moe', 'finetune'), + [ + pytest.param( + (1, 1, 1, 1, 1), + (1, 1, 1, 1, 1), + 'M*-M*-', + False, + False, + id='resume-without-finetune', + ), + pytest.param((1, 1, 1, 1, 1), (1, 1, 1, 1, 1), '*-*-', False, True, id='fsdp4-dense'), + pytest.param( + (2, 1, 1, 1, 1), (1, 1, 1, 1, 1), 'M*-M*-', False, True, id='tp2-fsdp2-to-fsdp4' + ), + pytest.param( + (1, 1, 2, 1, 1), (2, 1, 1, 1, 1), 'M*-M*-', False, True, id='cp2-fsdp2-to-tp2-fsdp2' + ), + pytest.param( + (1, 1, 1, 2, 1), + (2, 1, 1, 2, 2), + 'M*EM*E', + True, + True, + id='fsdp4-ep2-to-tp2-fsdp2-ep2-etp2', + ), + ], + ) + def test_gpt_fsdp_model_and_optimizer_load_into_hybrid( + self, tmp_path_dist_ckpt, src_parallel, dest_parallel, pattern, moe, finetune + ): + _run_gpt_to_hybrid_optimizer_load( + tmp_path_dist_ckpt, + src_parallel, + dest_parallel, + pattern, + moe, + finetune=finetune, + use_megatron_fsdp=True, + ) diff --git a/tests/unit_tests/dist_checkpointing/models/test_moe_experts.py b/tests/unit_tests/dist_checkpointing/models/test_moe_experts.py index 57de698ddff..126968a3af9 100644 --- a/tests/unit_tests/dist_checkpointing/models/test_moe_experts.py +++ b/tests/unit_tests/dist_checkpointing/models/test_moe_experts.py @@ -346,6 +346,7 @@ def test_sequential_grouped_mlp_interchangeable( def test_sequential_grouped_mlp_extra_state( self, tmp_path_dist_ckpt, + monkeypatch, src_tp_pp_exp, dest_tp_pp_exp, src_module, @@ -393,6 +394,8 @@ def test_sequential_grouped_mlp_extra_state( ckpt_dir_A, load_strategy, ) + # This checkpoint was created by the test and is therefore trusted. + monkeypatch.setenv("NVTE_ALLOW_UNSAFE_PICKLE_EXTRA_STATE", "1") model_A.load_state_dict( {k.removeprefix(layer_prefix): v for k, v in state_dict.items()} ) diff --git a/tests/unit_tests/dist_checkpointing/test_optimizer.py b/tests/unit_tests/dist_checkpointing/test_optimizer.py index 3849e044545..6b12ab04b24 100644 --- a/tests/unit_tests/dist_checkpointing/test_optimizer.py +++ b/tests/unit_tests/dist_checkpointing/test_optimizer.py @@ -25,7 +25,13 @@ get_gpt_layer_with_transformer_engine_spec as gpt_te_spec, ) from megatron.core.models.gpt.gpt_model import GPTModel -from megatron.core.optimizer import ChainedOptimizer, OptimizerConfig, get_megatron_optimizer +from megatron.core.optimizer import ( + HAVE_EMERGING_OPTIMIZERS, + ChainedOptimizer, + OptimizerConfig, + get_megatron_optimizer, +) +from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer from megatron.core.tensor_parallel import model_parallel_cuda_manual_seed from megatron.core.transformer import MLATransformerConfig, TransformerConfig from megatron.core.transformer.mlp import apply_swiglu_sharded_factory @@ -1135,6 +1141,99 @@ def test_model_parallel_dp_group_idx_preservation(self, tp, src_pp, dest_pp): # Check each dst group has at least 1 rank both in src and dest assert same_groups == set(range(num_dest_dp_groups)) + @pytest.mark.skipif( + not HAVE_EMERGING_OPTIMIZERS, reason="emerging_optimizers package not installed" + ) + @pytest.mark.skipif( + not is_torch_min_version("2.6a0"), reason="dp_reshardable requires PyTorch 2.6a0 or later" + ) + @pytest.mark.parametrize('sharding_type', ['dp_reshardable', 'fully_reshardable']) + def test_lion_optimizer_checkpoint_round_trip(self, tmp_path_dist_ckpt, sharding_type): + """Test DistributedOptimizer checkpoint save/load with Lion (single-moment optimizer). + + Lion is used as the scalar optimizer for Muon (muon_scalar_optimizer='lion'), + which is the natural path where Lion ends up inside a DistributedOptimizer. + This exercises the dynamic optimizer_state_keys logic with Lion's single + moment ('exp_avg') instead of Adam's two ('exp_avg', 'exp_avg_sq'). + """ + Utils.initialize_model_parallel(2, 1, order='tp-pp-dp') + + def _get_lion_distopt(optimizer): + """Extract the Lion DistributedOptimizer from a Muon+Lion ChainedOptimizer.""" + assert isinstance(optimizer, ChainedOptimizer) + for child in optimizer.chained_optimizers: + if isinstance(child, DistributedOptimizer): + return child + raise AssertionError("No DistributedOptimizer found in ChainedOptimizer") + + def _seed_random_optimizer_state(distopt, seed): + """Seed non-zero random exp_avg values in the DistOpt's raw optimizer state.""" + torch.manual_seed(seed) + for group in distopt.optimizer.param_groups: + for p in group['params']: + state = distopt.optimizer.state[p] + if 'exp_avg' in state: + state['exp_avg'].copy_(torch.randn_like(p.data)) + + with TempNamedDir( + tmp_path_dist_ckpt / 'test_lion_optimizer_checkpoint', sync=True + ) as ckpt_dir_A: + model_A, optimizer_A = setup_model_and_optimizer( + seed=2, + tp=2, + pp=1, + bf16=True, + dist_opt=True, + optimizer='muon', + muon_scalar_optimizer='lion', + use_param_layout=True, + ) + + lion_distopt_A = _get_lion_distopt(optimizer_A) + assert lion_distopt_A.optimizer_state_keys == ("exp_avg",) + _seed_random_optimizer_state(lion_distopt_A, seed=100) + + metadata = {'distrib_optim_sharding_type': sharding_type} + + model_sharded_sd = model_A[0].sharded_state_dict() + optim_sd = optimizer_A.sharded_state_dict(model_sharded_sd, metadata=metadata) + save(optim_sd, ckpt_dir_A) + + dp_zero_optim_A = lion_distopt_A.get_parameter_state_dp_zero(use_gloo_comm=False) + + model_B, optimizer_B = setup_model_and_optimizer( + seed=3, + tp=2, + pp=1, + bf16=True, + dist_opt=True, + optimizer='muon', + muon_scalar_optimizer='lion', + use_param_layout=True, + ) + + lion_distopt_B = _get_lion_distopt(optimizer_B) + _seed_random_optimizer_state(lion_distopt_B, seed=200) + + # Before loading, state should differ. + dp_zero_optim_B = lion_distopt_B.get_parameter_state_dp_zero(use_gloo_comm=False) + assert not self.check_equal_dp_zero_state(dp_zero_optim_A, dp_zero_optim_B, True) + + model_sharded_sd = model_B[0].sharded_state_dict() + load_sharded_state_dict = optimizer_B.sharded_state_dict( + model_sharded_sd, metadata=metadata, is_loading=True + ) + state_dict = load(load_sharded_state_dict, ckpt_dir_A) + optimizer_B.load_state_dict(state_dict) + + # After loading, state should match. + dp_zero_optim_B = lion_distopt_B.get_parameter_state_dp_zero(use_gloo_comm=False) + assert self.check_equal_dp_zero_state( + dp_zero_optim_A, dp_zero_optim_B, True, raise_if_different=True + ) + + Utils.destroy_model_parallel() + class TestFP32Optimizer: def setup_method(self, method): diff --git a/tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py b/tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py index 4ed91aa2cb6..aa1e682b39b 100644 --- a/tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py +++ b/tests/unit_tests/dist_checkpointing/test_pipeline_parallel_layout.py @@ -143,6 +143,7 @@ def create_args(): args.vocab_file = None args.add_position_embedding = False args.ckpt_assume_constant_structure = True + args.stream_ckpt_dequant = True args.ckpt_load_validate_sharding_integrity = True args.dist_ckpt_strictness = "assume_ok_unexpected" args.fp16 = False diff --git a/tests/unit_tests/dist_checkpointing/test_serialization.py b/tests/unit_tests/dist_checkpointing/test_serialization.py index 36de2e3c2c5..aadda4d5a06 100644 --- a/tests/unit_tests/dist_checkpointing/test_serialization.py +++ b/tests/unit_tests/dist_checkpointing/test_serialization.py @@ -1,4 +1,4 @@ -# Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import io import logging @@ -30,10 +30,15 @@ from megatron.core.dist_checkpointing.dict_utils import diff from megatron.core.dist_checkpointing.mapping import ShardedObject, ShardedTensorFactory from megatron.core.dist_checkpointing.serialization import ( + get_default_load_sharded_strategy, + get_default_save_sharded_strategy, load_sharded_metadata, load_tensors_metadata, ) -from megatron.core.dist_checkpointing.strategies.torch import TorchDistSaveShardedStrategy +from megatron.core.dist_checkpointing.strategies.torch import ( + TorchDistLoadShardedStrategy, + TorchDistSaveShardedStrategy, +) from megatron.core.dist_checkpointing.validation import StrictHandling from megatron.core.utils import is_torch_min_version from tests.unit_tests.dist_checkpointing import TempNamedDir @@ -47,6 +52,12 @@ def setup_method(self, method): def teardown_method(self, method): Utils.destroy_model_parallel() + def test_default_torch_dist_strategies(self): + assert isinstance(get_default_load_sharded_strategy(), TorchDistLoadShardedStrategy) + assert isinstance( + get_default_save_sharded_strategy("torch_dist"), TorchDistSaveShardedStrategy + ) + def test_single_process_save_load(self, tmp_path_dist_ckpt): Utils.initialize_model_parallel(1, 1) @@ -528,8 +539,6 @@ def test_tensor_shape_mismatch(self, tmp_path_dist_ckpt): not is_torch_min_version("2.3.0"), reason="remove_sharded_tensors relies on Torch APIs introduced in v2.3.0", ) - @pytest.mark.flaky - @pytest.mark.flaky_in_dev def test_remove_sharded_tensors(self, tmp_path_dist_ckpt): Utils.initialize_model_parallel(2, 4) @@ -576,7 +585,10 @@ def test_remove_sharded_tensors(self, tmp_path_dist_ckpt): assert len(prefix_files) == 0 new_metadata = fs_reader.read_metadata() - assert set(new_metadata.state_dict_metadata.keys()) == {'keyA'} + assert set(new_metadata.state_dict_metadata.keys()) == { + 'common_state/shard_0_1', + 'keyA', + } Utils.destroy_model_parallel() diff --git a/tests/unit_tests/dist_checkpointing/test_stream_ckpt_dequant.py b/tests/unit_tests/dist_checkpointing/test_stream_ckpt_dequant.py new file mode 100644 index 00000000000..bc940668469 --- /dev/null +++ b/tests/unit_tests/dist_checkpointing/test_stream_ckpt_dequant.py @@ -0,0 +1,388 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Tests for the streaming per-tensor dequantize path used when loading +distributed checkpoints with quantized (FP8 / MXFP8 / blockwise / NVFP4) +model parameters. + +The feature under test is ``stream_ckpt_dequant`` on ``MCoreLoadPlanner`` and +``TorchDistLoadShardedStrategy``. When on, the LoadPlanner dequantizes each +quantized destination one at a time inside ``resolve_tensor``/``commit_tensor`` +instead of up-front in ``force_all_tensors_to_non_fp8``. Tests cover: + +- Loaded-content equivalence vs. the legacy upfront path (FP8). +- Delayed-scaling ``amax_history`` is not polluted across a streaming load. +- ``_unwrap_pyt_sharded_tensor`` uses view-based axis stripping (no + dequantize fallback) — exercised implicitly by the MXFP8 save/load test. +- MXFP8 save/load round-trip. +- NVFP4 save/load round-trip (Blackwell+ only, skipped otherwise). +- No-op fall-through for plain (non-quantized) tensors. +""" + +import pytest +import torch + +try: + from transformer_engine.pytorch.float8_tensor import Float8Tensor + from transformer_engine.pytorch.tensor import QuantizedTensor + + HAVE_TE = True +except ImportError: + HAVE_TE = False + Float8Tensor = None # type: ignore + QuantizedTensor = None # type: ignore + +try: + import transformer_engine.pytorch.tensor.mxfp8_tensor # noqa: F401 + + HAVE_MXFP8 = True +except ImportError: + HAVE_MXFP8 = False + +try: + import transformer_engine.pytorch.tensor.nvfp4_tensor # noqa: F401 + + HAVE_NVFP4 = True +except ImportError: + HAVE_NVFP4 = False + +try: + from megatron.training.utils import get_device_arch_version + + _DEVICE_ARCH = get_device_arch_version() +except Exception: + _DEVICE_ARCH = 0 + +# MXFP8 and NVFP4 require Blackwell (arch 10+). +HAVE_MXFP8_HW = HAVE_MXFP8 and _DEVICE_ARCH >= 10 +HAVE_NVFP4_HW = HAVE_NVFP4 and _DEVICE_ARCH >= 10 + +from megatron.core.dist_checkpointing import ShardedTensor, load, save +from megatron.core.dist_checkpointing.strategies.torch import ( + MCoreLoadPlanner, + TorchDistLoadShardedStrategy, + TorchDistSaveShardedStrategy, +) +from tests.unit_tests.dist_checkpointing import TempNamedDir +from tests.unit_tests.test_utilities import Utils + + +def _to_float8(tensor: torch.Tensor): + """Convert a BF16 tensor to delayed-scaling Float8Tensor (TE 2.x API).""" + try: + return Float8Tensor.to_float8(tensor) + except Exception: + import transformer_engine_torch as tex + from transformer_engine.pytorch.tensor.float8_tensor import Float8Quantizer + + quantizer = Float8Quantizer( + scale=torch.full([1], 1.0, dtype=torch.float32, device="cuda"), + amax=torch.empty([1], dtype=torch.float32, device="cuda"), + fp8_dtype=tex.DType.kFloat8E4M3, + ) + return quantizer(tensor.cuda()) + + +def _to_mxfp8(tensor: torch.Tensor): + """Convert a BF16 tensor to MXFP8Tensor.""" + import transformer_engine_torch as tex + from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Quantizer + + quantizer = MXFP8Quantizer(fp8_dtype=tex.DType.kFloat8E4M3) + return quantizer(tensor.cuda().contiguous()) + + +def _to_nvfp4(tensor: torch.Tensor): + """Convert a BF16 tensor to NVFP4Tensor (Blackwell+ only).""" + from transformer_engine.pytorch.tensor.nvfp4_tensor import NVFP4Quantizer + + quantizer = NVFP4Quantizer( + rowwise=True, + columnwise=True, + with_rht=False, + with_post_rht_amax=False, + with_2d_quantization=True, + stochastic_rounding=False, + with_random_sign_mask=False, + ) + return quantizer(tensor.cuda().contiguous()) + + +@pytest.mark.skipif(not HAVE_TE, reason="TransformerEngine not available") +class TestStreamCkptDequant: + """Unit tests for streaming per-tensor dequantize during ckpt load.""" + + def teardown_method(self, method): + Utils.destroy_model_parallel() + + # --------------------------------------------------------------- + # Baseline: FP8 (delayed scaling) save/load equivalence + amax safety + # --------------------------------------------------------------- + + @pytest.mark.parametrize('stream_ckpt_dequant', [False, True]) + def test_fp8_save_load_content_equivalence(self, tmp_path_dist_ckpt, stream_ckpt_dequant): + """Loaded FP8 contents must match regardless of which dequantize path is used.""" + Utils.initialize_model_parallel(1, 1) + + fill_val = 0.5 + + def get_fp8_tensor(val): + return _to_float8(torch.full((8,), val, dtype=torch.bfloat16, device='cuda')) + + def get_state_dict(val): + return { + 'w': ShardedTensor.from_rank_offsets( + 'w', get_fp8_tensor(val), replica_id=Utils.rank + ) + } + + with TempNamedDir(tmp_path_dist_ckpt / f'fp8_eq_{stream_ckpt_dequant}') as ckpt_dir: + save(get_state_dict(fill_val), ckpt_dir, TorchDistSaveShardedStrategy()) + + # Fresh state dict with a different fill — the load must overwrite it. + sd_to_load = get_state_dict(99.0) + strategy = TorchDistLoadShardedStrategy(stream_ckpt_dequant=stream_ckpt_dequant) + loaded = load(sd_to_load, ckpt_dir, strategy) + # Dequantize the loaded tensor (may be Float8 or BF16 depending on path) + loaded_w = loaded['w'] + if isinstance(loaded_w, QuantizedTensor): + loaded_w = loaded_w.dequantize() + # fill_val (0.5) is exactly representable in FP8 E4M3 and the per-tensor + # scale is a power of 2, so the round-trip is numerically lossless modulo + # bf16 rounding. Tight tolerance catches real regressions. + torch.testing.assert_close( + loaded_w, + torch.full((8,), fill_val, dtype=torch.bfloat16, device='cuda'), + rtol=1e-3, + atol=1e-3, + ) + + def test_fp8_amax_history_not_polluted(self, tmp_path_dist_ckpt): + """Delayed-scaling amax must be snapshotted & restored across a streaming load.""" + Utils.initialize_model_parallel(1, 1) + + def get_fp8_tensor(val): + return _to_float8(torch.full((8,), val, dtype=torch.bfloat16, device='cuda')) + + sd_to_save = { + 'w': ShardedTensor.from_rank_offsets('w', get_fp8_tensor(0.25), replica_id=Utils.rank) + } + + with TempNamedDir(tmp_path_dist_ckpt / 'fp8_amax') as ckpt_dir: + save(sd_to_save, ckpt_dir, TorchDistSaveShardedStrategy()) + + # Rebuild destination with a known (distinct) amax value we can check + # survives the streaming load. + dst = ShardedTensor.from_rank_offsets('w', get_fp8_tensor(99.0), replica_id=Utils.rank) + q = getattr(dst.data, "_quantizer", None) + if q is None or not isinstance(getattr(q, "amax", None), torch.Tensor): + pytest.skip("This TE build's Float8Tensor has no quantizer.amax scalar") + sentinel = 42.0 + q.amax.fill_(sentinel) + pre_load_amax = q.amax.detach().clone() + + strategy = TorchDistLoadShardedStrategy(stream_ckpt_dequant=True) + loaded = load({'w': dst}, ckpt_dir, strategy) + + # The loaded tensor is the same QuantizedTensor object; its quantizer.amax + # must be exactly what we put in before the load. + loaded_q = getattr(loaded['w'], "_quantizer", None) + assert loaded_q is not None + assert torch.equal(loaded_q.amax, pre_load_amax), ( + f"amax was not restored after streaming load; " + f"before={pre_load_amax.item()} after={loaded_q.amax.item()}" + ) + + # --------------------------------------------------------------- + # MXFP8: exercises both the streaming dequant AND the view-based + # _unwrap_pyt_sharded_tensor fix (without it, ten[0] on MXFP8 OOMs). + # --------------------------------------------------------------- + + @pytest.mark.skipif( + not HAVE_MXFP8_HW, + reason="MXFP8 requires TransformerEngine MXFP8Tensor and Blackwell+ (arch 10+)", + ) + @pytest.mark.parametrize('stream_ckpt_dequant', [False, True]) + def test_mxfp8_save_load_content_equivalence(self, tmp_path_dist_ckpt, stream_ckpt_dequant): + Utils.initialize_model_parallel(1, 1) + + # MXFP8 requires 2D with last-dim aligned to block size (32). + fill_val = 0.25 + + def get_mxfp8_tensor(val): + return _to_mxfp8(torch.full((64, 128), val, dtype=torch.bfloat16, device='cuda')) + + def get_state_dict(val): + return { + 'w': ShardedTensor.from_rank_offsets( + 'w', get_mxfp8_tensor(val), replica_id=Utils.rank + ) + } + + with TempNamedDir(tmp_path_dist_ckpt / f'mxfp8_eq_{stream_ckpt_dequant}') as ckpt_dir: + save(get_state_dict(fill_val), ckpt_dir, TorchDistSaveShardedStrategy()) + + sd_to_load = get_state_dict(99.0) + strategy = TorchDistLoadShardedStrategy(stream_ckpt_dequant=stream_ckpt_dequant) + loaded = load(sd_to_load, ckpt_dir, strategy) + loaded_w = loaded['w'] + if isinstance(loaded_w, QuantizedTensor): + loaded_w = loaded_w.dequantize() + # fill_val (0.25) is exactly representable in FP8 E4M3, and MXFP8 stores + # block scales in E8M0 (power-of-2), so the per-block scale is exact and + # every element encodes to the same FP8 code. Round-trip is near-lossless. + torch.testing.assert_close( + loaded_w, + torch.full((64, 128), fill_val, dtype=torch.bfloat16, device='cuda'), + rtol=1e-3, + atol=1e-3, + ) + + # --------------------------------------------------------------- + # NVFP4: round-trip under both paths. Same invariants as MXFP8 but + # with NVFP4Tensor — validates that `is_float8tensor` (which binds to + # QuantizedTensor under TE 2.x) correctly covers the FP4 path, that + # NVFP4Tensor.view works inside _unwrap_pyt_sharded_tensor, and that + # BF16->NVFP4 copy through QuantizedTensor.__torch_dispatch__ -> quantize_ + # produces correct values. Requires Blackwell+ for the FP4 kernels. + # --------------------------------------------------------------- + + @pytest.mark.skipif( + not HAVE_NVFP4_HW, + reason="NVFP4 requires TransformerEngine NVFP4Tensor and Blackwell+ (arch 10+)", + ) + @pytest.mark.parametrize('stream_ckpt_dequant', [False, True]) + def test_nvfp4_save_load_content_equivalence(self, tmp_path_dist_ckpt, stream_ckpt_dequant): + Utils.initialize_model_parallel(1, 1) + + # NVFP4BlockScaling uses 16-element blocks along the last dim; use a + # shape that's a multiple of both common block sizes. + fill_val = 0.25 + + def get_nvfp4_tensor(val): + return _to_nvfp4(torch.full((64, 128), val, dtype=torch.bfloat16, device='cuda')) + + def get_state_dict(val): + return { + 'w': ShardedTensor.from_rank_offsets( + 'w', get_nvfp4_tensor(val), replica_id=Utils.rank + ) + } + + with TempNamedDir(tmp_path_dist_ckpt / f'nvfp4_eq_{stream_ckpt_dequant}') as ckpt_dir: + save(get_state_dict(fill_val), ckpt_dir, TorchDistSaveShardedStrategy()) + + sd_to_load = get_state_dict(99.0) + strategy = TorchDistLoadShardedStrategy(stream_ckpt_dequant=stream_ckpt_dequant) + loaded = load(sd_to_load, ckpt_dir, strategy) + loaded_w = loaded['w'] + if isinstance(loaded_w, QuantizedTensor): + loaded_w = loaded_w.dequantize() + # For a constant block every FP4 code is identical and the dominant error + # source is the per-block scale being stored in FP8 E4M3 (unlike MXFP8's + # power-of-2 E8M0). That rounding is bounded below ~1% relative; 1e-2 is + # tight enough to catch real bugs and loose enough to absorb E4M3 scale + # rounding + bf16 output rounding. + torch.testing.assert_close( + loaded_w, + torch.full((64, 128), fill_val, dtype=torch.bfloat16, device='cuda'), + rtol=1e-2, + atol=1e-2, + ) + + # --------------------------------------------------------------- + # Corner cases + # --------------------------------------------------------------- + + @pytest.mark.parametrize('stream_ckpt_dequant', [False, True]) + def test_plain_tensor_untouched_by_streaming_path( + self, tmp_path_dist_ckpt, stream_ckpt_dequant + ): + """Non-quantized tensors in the state dict must round-trip losslessly under either path.""" + Utils.initialize_model_parallel(1, 1) + + src = torch.arange(64, dtype=torch.bfloat16, device='cuda') + sd_to_save = {'w': ShardedTensor.from_rank_offsets('w', src.clone(), replica_id=Utils.rank)} + + with TempNamedDir(tmp_path_dist_ckpt / f'plain_{stream_ckpt_dequant}') as ckpt_dir: + save(sd_to_save, ckpt_dir, TorchDistSaveShardedStrategy()) + + dst = { + 'w': ShardedTensor.from_rank_offsets( + 'w', torch.zeros_like(src), replica_id=Utils.rank + ) + } + strategy = TorchDistLoadShardedStrategy(stream_ckpt_dequant=stream_ckpt_dequant) + loaded = load(dst, ckpt_dir, strategy) + # Plain BF16 must round-trip exactly. + torch.testing.assert_close(loaded['w'], src) + + def test_default_is_on(self): + """The default for stream_ckpt_dequant must be True (streaming path is now default).""" + strat = TorchDistLoadShardedStrategy() + assert ( + strat.stream_ckpt_dequant is True + ), "Default must be True; users opt out via --no-stream-ckpt-dequant." + planner = MCoreLoadPlanner() + assert planner.stream_ckpt_dequant is True + + def test_planner_state_cleanup_after_load(self, tmp_path_dist_ckpt): + """``_intermediate_read_items`` must be empty after a streaming load completes. + + Lingering entries would indicate a scratch tensor we forgot to drop, defeating + the memory win. + """ + Utils.initialize_model_parallel(1, 1) + + def get_fp8_tensor(val): + return _to_float8(torch.full((32,), val, dtype=torch.bfloat16, device='cuda')) + + sd_to_save = { + f'w{i}': ShardedTensor.from_rank_offsets( + f'w{i}', get_fp8_tensor(0.125), replica_id=Utils.rank + ) + for i in range(4) + } + + with TempNamedDir(tmp_path_dist_ckpt / 'planner_cleanup') as ckpt_dir: + save(sd_to_save, ckpt_dir, TorchDistSaveShardedStrategy()) + + # Instrument: intercept MCoreLoadPlanner to capture the live instance. + captured: list[MCoreLoadPlanner] = [] + original_init = MCoreLoadPlanner.__init__ + + def capturing_init(self, *args, **kwargs): + original_init(self, *args, **kwargs) + captured.append(self) + + MCoreLoadPlanner.__init__ = capturing_init # type: ignore[assignment] + try: + dst = { + f'w{i}': ShardedTensor.from_rank_offsets( + f'w{i}', get_fp8_tensor(99.0), replica_id=Utils.rank + ) + for i in range(4) + } + load(dst, ckpt_dir, TorchDistLoadShardedStrategy(stream_ckpt_dequant=True)) + finally: + MCoreLoadPlanner.__init__ = original_init # type: ignore[assignment] + + assert len(captured) == 1 + assert captured[0]._intermediate_read_items == {}, ( + f"Planner left intermediate state after load: " + f"{list(captured[0]._intermediate_read_items.keys())}" + ) + + def test_streaming_flag_forwards_through_fpsl_wrapper(self): + """FullyParallelLoadStrategyWrapper must surface the base strategy's flag.""" + from megatron.core.dist_checkpointing.strategies.fully_parallel import ( + FullyParallelLoadStrategyWrapper, + ) + + base_off = TorchDistLoadShardedStrategy(stream_ckpt_dequant=False) + base_on = TorchDistLoadShardedStrategy(stream_ckpt_dequant=True) + # parallelization_group left default -> GroupMember.WORLD; that's fine since + # we're only reading the forwarded property, not calling load(). + wrapped_off = FullyParallelLoadStrategyWrapper(base_off) + wrapped_on = FullyParallelLoadStrategyWrapper(base_on) + assert wrapped_off.stream_ckpt_dequant is False + assert wrapped_on.stream_ckpt_dequant is True diff --git a/tests/unit_tests/dist_checkpointing/test_validation.py b/tests/unit_tests/dist_checkpointing/test_validation.py new file mode 100644 index 00000000000..80313ee8353 --- /dev/null +++ b/tests/unit_tests/dist_checkpointing/test_validation.py @@ -0,0 +1,27 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +from unittest.mock import Mock + +import torch + +from megatron.core.dist_checkpointing.validation import determine_global_metadata + + +def test_determine_global_metadata_uses_explicit_process_group(monkeypatch): + process_group = Mock() + metadata = Mock() + shard = Mock() + shard.without_data.return_value = metadata + get_world_size = Mock(return_value=2) + all_gather_object = Mock() + monkeypatch.setattr(torch.distributed, "get_world_size", get_world_size) + monkeypatch.setattr(torch.distributed, "all_gather_object", all_gather_object) + + local_metadata, global_metadata = determine_global_metadata( + {"model": shard}, process_group=process_group + ) + + assert local_metadata == [metadata] + assert global_metadata == [None, None] + get_world_size.assert_called_once_with(group=process_group) + all_gather_object.assert_called_once_with(global_metadata, local_metadata, group=process_group) diff --git a/tests/unit_tests/dist_checkpointing/utils.py b/tests/unit_tests/dist_checkpointing/utils.py index 81851f7f80e..9a8b502eba0 100644 --- a/tests/unit_tests/dist_checkpointing/utils.py +++ b/tests/unit_tests/dist_checkpointing/utils.py @@ -150,6 +150,7 @@ def init_checkpointing_mock_args(args, ckpt_dir, fully_parallel=False): args.no_save_optim = False args.no_save_rng = False args.ckpt_assume_constant_structure = False + args.stream_ckpt_dequant = True args.ckpt_load_validate_sharding_integrity = True args.log_progress = False args.auto_detect_ckpt_format = False @@ -185,6 +186,11 @@ def setup_model_and_optimizer( dist_opt=True, optimizer='adam', use_param_layout=False, + muon_scalar_optimizer='adam', + cp=1, + ep=1, + etp=1, + use_megatron_fsdp=False, ): optimizer_type = optimizer use_layer_wise = False @@ -205,6 +211,20 @@ def setup_model_and_optimizer( mock_args = parse_args(ignore_unknown_args=True) with mock.patch('megatron.training.training.get_args', new=lambda: mock_args): init_basic_mock_args(mock_args, tp, pp, bf16=bf16) + mock_args.context_parallel_size = cp + mock_args.expert_model_parallel_size = ep + mock_args.expert_tensor_parallel_size = etp + mock_args.use_megatron_fsdp = use_megatron_fsdp + mock_args.data_parallel_sharding_strategy = ( + 'optim_grads_params' if use_megatron_fsdp else 'no_shard' + ) + if use_megatron_fsdp: + # parse_args() leaves these as CLI strings until validate_args() + # maps them to the torch.dtype values expected by Megatron-FSDP. + mock_args.megatron_fsdp_main_params_dtype = torch.float32 + mock_args.megatron_fsdp_main_grads_dtype = None + mock_args.megatron_fsdp_grad_comm_dtype = None + mock_args.gradient_accumulation_fusion = False mock_args.use_distributed_optimizer = ddp_use_dist_opt mock_args.use_layer_wise_distributed_optimizer = ddp_use_layer_wise if ddp_use_layer_wise: @@ -216,6 +236,9 @@ def setup_model_and_optimizer( tensor_model_parallel_size=tp, pipeline_model_parallel_size=pp, pipeline_dtype=torch.bfloat16, + context_parallel_size=cp, + expert_model_parallel_size=ep, + expert_tensor_parallel_size=etp, bf16=bf16, ) ) @@ -226,10 +249,17 @@ def setup_model_and_optimizer( use_distributed_optimizer=ddp_use_dist_opt, use_layer_wise_distributed_optimizer=use_layer_wise, optimizer=optimizer, + muon_scalar_optimizer=muon_scalar_optimizer, ) + if use_megatron_fsdp: + # The FSDP DTensor sharded-state path may materialize missing optimizer + # slots with a dummy step, which requires a concrete learning rate. + config.lr = 1.0e-3 if optimizer_type in ('muon', 'dist_muon'): config.lr = 0.0 + elif optimizer_type == 'lion': + config.lr = 1e-4 optimizer = get_megatron_optimizer(config, model) torch.manual_seed(seed + 1) @@ -255,13 +285,20 @@ def _init_states(optimizer): if isinstance(optimizer, ChainedOptimizer): _init_states(optimizer) else: + if hasattr(optimizer, 'optimizer_state_keys'): + state_keys = optimizer.optimizer_state_keys + else: + state_keys = ("exp_avg", "exp_avg_sq") for group in optimizer.optimizer.param_groups: for p in group['params']: if len(optimizer.optimizer.state[p]) == 0: - optimizer.optimizer.state[p]['exp_avg'] = torch.rand_like(p.data) - optimizer.optimizer.state[p]['exp_avg_sq'] = torch.rand_like(p.data) + for key in state_keys: + optimizer.optimizer.state[p][key] = torch.rand_like(p.data) - optimizer.reload_model_params() + # Megatron-FSDP owns the model/main-parameter synchronization and its + # DistributedOptimizer intentionally does not implement this legacy copy. + if not use_megatron_fsdp: + optimizer.reload_model_params() CachedMetadataFileSystemReader.clear_metadata_cache() return unwrap_model(model), optimizer diff --git a/tests/unit_tests/distributed/mfsdp_v1/test_mfsdp_fully_shard.py b/tests/unit_tests/distributed/mfsdp_v1/test_mfsdp_fully_shard.py index 2bc198695b0..7b3a8fd9c9a 100644 --- a/tests/unit_tests/distributed/mfsdp_v1/test_mfsdp_fully_shard.py +++ b/tests/unit_tests/distributed/mfsdp_v1/test_mfsdp_fully_shard.py @@ -136,6 +136,23 @@ def forward(self, x, y): return x +class RootParamModel(torch.nn.Module): + """Toy model with parameters owned directly by the root module.""" + + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter(torch.empty(DIM_SIZE, DIM_SIZE)) + self.bias = torch.nn.Parameter(torch.empty(DIM_SIZE)) + self.reset_parameters() + + def reset_parameters(self): + torch.nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5)) + torch.nn.init.zeros_(self.bias) + + def forward(self, x): + return torch.nn.functional.linear(x, self.weight, self.bias) + + class ToyTETransformer(torch.nn.Module): """Toy Transformer model for testing Megatron-FSDP with Transformer Engine.""" @@ -731,6 +748,33 @@ def test_fully_shard_ez(self, shard_strategy): optimizer.step() optimizer.zero_grad() + def test_root_module_forward_uses_gathered_parameters(self): + """ + Test that root-owned parameters are gathered before the root forward. + """ + + model = RootParamModel().cuda() + with torch.no_grad(): + model.weight.copy_( + torch.arange(DIM_SIZE * DIM_SIZE, dtype=torch.float32, device="cuda").view( + DIM_SIZE, DIM_SIZE + ) + ) + model.bias.copy_(torch.arange(DIM_SIZE, dtype=torch.float32, device="cuda")) + + model_input = torch.arange(DIM_SIZE * DIM_SIZE, dtype=torch.float32, device="cuda").view( + DIM_SIZE, DIM_SIZE + ) + expected_output = model(model_input) + + mfsdp_model = fully_shard_model( + module=model, fsdp_unit_modules=[RootParamModel], zero_dp_strategy=OPTIM_GRADS_PARAMS + ) + + output = mfsdp_model(model_input) + + torch.testing.assert_close(output, expected_output) + @pytest.mark.skipif( version.parse(torch.__version__) < version.parse('2.4.0'), reason="Megatron-FSDP requires PyTorch 2.4.0 or later.", diff --git a/tests/unit_tests/distributed/mfsdp_v2/conftest.py b/tests/unit_tests/distributed/mfsdp_v2/conftest.py index cff48b29fce..741c0f57e82 100644 --- a/tests/unit_tests/distributed/mfsdp_v2/conftest.py +++ b/tests/unit_tests/distributed/mfsdp_v2/conftest.py @@ -21,6 +21,13 @@ class DistributedSetup: @pytest.fixture(scope="function") def distributed_setup() -> Iterator[DistributedSetup]: """Read torchrun rank state and set up this rank's local device.""" + # Some MFSDP v2 tests are sensitive to NCCL algorithm/channel choices. Clear + # the suite-wide NCCL defaults (set in the top-level conftest.py) before + # init_device_mesh initializes NCCL communicators so this bucket uses NCCL + # settings closer to production. + os.environ.pop("NCCL_MAX_NCHANNELS", None) + os.environ.pop("NCCL_NVLS_ENABLE", None) + if "RANK" not in os.environ or "WORLD_SIZE" not in os.environ: pytest.skip("Not running under torchrun. Use torchrun to run this test file.") diff --git a/tests/unit_tests/distributed/mfsdp_v2/profiler_utils.py b/tests/unit_tests/distributed/mfsdp_v2/profiler_utils.py new file mode 100644 index 00000000000..426b2db49dc --- /dev/null +++ b/tests/unit_tests/distributed/mfsdp_v2/profiler_utils.py @@ -0,0 +1,64 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Helpers for parsing ``torch.profiler`` events in the mfsdp_v2 tests.""" + +import pytest +from torch.autograd import DeviceType +from torch.autograd.profiler_util import FunctionEvent +from torch.profiler import profile as TorchProfiler + + +def events_overlap(first: FunctionEvent, second: FunctionEvent) -> bool: + return ( + first.time_range.start < second.time_range.end + and second.time_range.start < first.time_range.end + ) + + +def collect_linked_kernels( + prof: TorchProfiler, cpu_event_name_substring: str +) -> list[FunctionEvent]: + """Collect device kernel events linked to matching CPU op instances. + + Device events are attributed by their launching CPU op rather than searched by their + own name: device-side names vary across GPU architectures and kernel libraries -- for + example a matmul kernel is named ``nvjet_``/``cutlass_``/``cublas_``... while its + CPU op is simply ``aten::mm``. + + Zero-CTA all-gather copy-engine memcpys are not kernels and are intentionally not + returned. + """ + # A correlation id is shared by a device event and the leaf runtime op that issued it, + # not the enclosing matched op, so walk cpu_parent up from each correlated leaf. Id 0 + # is the "no device correlation" sentinel and is skipped. + events = prof.events() + # ``FunctionEvent.linked_correlation_id`` was added to the torch profiler API after the + # PyTorch release this container is pinned to; the attribute-based kernel linking below + # is unavailable there. Skip rather than fail so these overlap tests activate automatically + # once the base image advances to a torch that exposes it. + if events and not hasattr(events[0], "linked_correlation_id"): + pytest.skip( + "torch.profiler FunctionEvent lacks 'linked_correlation_id' in this torch build" + ) + matching_correlations: set[int] = set() + for event in events: + if event.device_type != DeviceType.CPU or not event.linked_correlation_id: + continue + node = event + while node is not None: + if cpu_event_name_substring in node.name: + matching_correlations.add(event.linked_correlation_id) + break + node = node.cpu_parent + + linked_kernels: list[FunctionEvent] = [] + for event in events: + if event.device_type != DeviceType.CUDA: + continue + if event.activity_type != "kernel": + continue + if event.linked_correlation_id not in matching_correlations: + continue + linked_kernels.append(event) + + return linked_kernels diff --git a/tests/unit_tests/distributed/mfsdp_v1/test_annotation.py b/tests/unit_tests/distributed/mfsdp_v2/test_annotation.py similarity index 59% rename from tests/unit_tests/distributed/mfsdp_v1/test_annotation.py rename to tests/unit_tests/distributed/mfsdp_v2/test_annotation.py index 9aee7734172..cdb6db0f182 100644 --- a/tests/unit_tests/distributed/mfsdp_v1/test_annotation.py +++ b/tests/unit_tests/distributed/mfsdp_v2/test_annotation.py @@ -38,6 +38,17 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: return x +class FrozenFirstLayerModel(nn.Module): + def __init__(self, dim: int) -> None: + super().__init__() + self.bias = nn.Parameter(torch.ones(dim)) + self.layers = nn.ModuleList([nn.Linear(dim, dim, bias=False) for _ in range(2)]) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = torch.relu(self.layers[0](x)) + return torch.relu(self.layers[1](x + self.bias)) + + def _flat_placements() -> Placements: return Placements(dp_axes=[0], parameter=[Flat()], gradient=[Flat()], optimizer=[Flat()]) @@ -63,16 +74,10 @@ def record_pop() -> None: monkeypatch.setattr(torch.cuda.nvtx, "range_pop", record_pop) -def _get_distributed_setup(request: pytest.FixtureRequest): - try: - return request.getfixturevalue("distributed_setup") - except pytest.FixtureLookupError: - pytest.skip("distributed_setup fixture is only available in the Megatron-FSDP test bucket") - - -def test_fsdp_sibling_roots_emit_root_nvtx_ranges_after_training_step(request, monkeypatch): +def test_fsdp_sibling_roots_emit_root_nvtx_ranges_after_training_step( + distributed_setup, monkeypatch +): """Independent FSDP roots should each emit root-labeled NVTX ranges.""" - distributed_setup = _get_distributed_setup(request) events: list[NvtxEvent] = [] _setup_nvtx_recording(monkeypatch, events) model = NestedLinearModel(dim=4).to(distributed_setup.device) @@ -94,9 +99,8 @@ def test_fsdp_sibling_roots_emit_root_nvtx_ranges_after_training_step(request, m ] -def test_fsdp_training_hooks_emit_stacked_nvtx_ranges(request, monkeypatch): +def test_fsdp_training_hooks_emit_stacked_nvtx_ranges(distributed_setup, monkeypatch): """Nested training hooks should emit concise NVTX ranges.""" - distributed_setup = _get_distributed_setup(request) events: list[NvtxEvent] = [] _setup_nvtx_recording(monkeypatch, events) model = NestedLinearModel(dim=4).to(distributed_setup.device) @@ -121,3 +125,54 @@ def test_fsdp_training_hooks_emit_stacked_nvtx_ranges(request, monkeypatch): ("pop", "layers.0", "backward"), ("pop", "", "backward"), ] + + +def test_fsdp_frozen_parameters_emit_balanced_backward_nvtx_range(distributed_setup, monkeypatch): + """Frozen FSDP units should still balance backward NVTX ranges.""" + events: list[NvtxEvent] = [] + _setup_nvtx_recording(monkeypatch, events) + model = nn.Linear(4, 4, bias=False).to(distributed_setup.device) + for parameter in model.parameters(): + parameter.requires_grad_(False) + mesh = init_device_mesh(distributed_setup.device.type, (distributed_setup.world_size,)) + fully_shard(model, mesh=mesh, placements=_flat_placements()) + + x = torch.ones(2, 4, device=distributed_setup.device, requires_grad=True) + model(x).sum().backward() + + assert [(event.kind, event.name, event.phase) for event in events] == [ + ("push", "", "forward"), + ("pop", "", "forward"), + ("push", "", "backward"), + ("pop", "", "backward"), + ] + + +def test_fsdp_frozen_child_without_grad_inputs_skips_backward_nvtx_range( + distributed_setup, monkeypatch +): + """Frozen FSDP children outside the backward graph should not emit backward ranges.""" + events: list[NvtxEvent] = [] + _setup_nvtx_recording(monkeypatch, events) + model = FrozenFirstLayerModel(dim=4).to(distributed_setup.device) + for parameter in model.layers[0].parameters(): + parameter.requires_grad_(False) + mesh = init_device_mesh(distributed_setup.device.type, (distributed_setup.world_size,)) + fully_shard(model.layers[0], mesh=mesh, placements=_flat_placements()) + fully_shard(model.layers[1], mesh=mesh, placements=_flat_placements()) + fully_shard(model, mesh=mesh, placements=_flat_placements()) + + model(torch.ones(2, 4, device=distributed_setup.device)).sum().backward() + + assert [(event.kind, event.name, event.phase) for event in events] == [ + ("push", "", "forward"), + ("push", "layers.0", "forward"), + ("pop", "layers.0", "forward"), + ("push", "layers.1", "forward"), + ("pop", "layers.1", "forward"), + ("pop", "", "forward"), + ("push", "", "backward"), + ("push", "layers.1", "backward"), + ("pop", "layers.1", "backward"), + ("pop", "", "backward"), + ] diff --git a/tests/unit_tests/distributed/mfsdp_v2/test_context.py b/tests/unit_tests/distributed/mfsdp_v2/test_context.py index 9ffd6bd8cdb..102b2fde332 100644 --- a/tests/unit_tests/distributed/mfsdp_v2/test_context.py +++ b/tests/unit_tests/distributed/mfsdp_v2/test_context.py @@ -27,7 +27,7 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: class MultiChildModel(nn.Module): - """Model with direct parameters and multiple child FSDP units.""" + """Model with direct parameters and multiple child FsdpModules.""" def __init__(self, dim: int, num_children: int) -> None: super().__init__() @@ -42,12 +42,39 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: return x +class BranchModel(nn.Module): + """Nested branch with its own child FsdpModule.""" + + def __init__(self, dim: int) -> None: + super().__init__() + self.bias = nn.Parameter(torch.ones(dim)) + self.inner = nn.Linear(dim, dim, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Run the nested branch.""" + return torch.relu(self.inner(x) + self.bias) + + +class NestedSiblingModel(nn.Module): + """Model with a nested left subtree and a right sibling.""" + + def __init__(self, dim: int) -> None: + super().__init__() + self.bias = nn.Parameter(torch.ones(dim)) + self.left = BranchModel(dim) + self.right = nn.Linear(dim, dim, bias=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Run the nested subtree before the right sibling.""" + return self.right(self.left(x) + self.bias) + + def _flat_placements() -> Placements: return Placements(dp_axes=[0], parameter=[Flat()], gradient=[Flat()], optimizer=[Flat()]) def test_child_then_parent_share_one_context(distributed_setup): - """A parent FSDP unit should lazily create one context for its subtree.""" + """A parent FsdpModule should lazily create one context for its subtree.""" device = distributed_setup.device mesh = init_device_mesh(device.type, (distributed_setup.world_size,)) @@ -98,3 +125,23 @@ def test_sibling_roots_without_parent_keep_separate_contexts(distributed_setup): assert model.layers[0].context is not model.layers[1].context assert model.layers[0].is_root() assert model.layers[1].is_root() + + +def test_nested_prefetch_orders_use_dfs(distributed_setup): + """Nested FsdpModules should use DFS orders for one-step prefetch.""" + device = distributed_setup.device + + mesh = init_device_mesh(device.type, (distributed_setup.world_size,)) + model = NestedSiblingModel(dim=4).to(device) + + fully_shard(model.left.inner, mesh=mesh, placements=_flat_placements()) + fully_shard(model.left, mesh=mesh, placements=_flat_placements()) + fully_shard(model.right, mesh=mesh, placements=_flat_placements()) + fully_shard(model, mesh=mesh, placements=_flat_placements()) + + with torch.no_grad(): + model(torch.ones(2, 4, device=device)) + + context = model.context + assert list(context.forward_order) == [model, model.left, model.left.inner, model.right] + assert list(context.backward_order) == [model, model.right, model.left, model.left.inner] diff --git a/tests/unit_tests/distributed/mfsdp_v2/test_cuda_graph.py b/tests/unit_tests/distributed/mfsdp_v2/test_cuda_graph.py index 910c13c6fd3..08920bcea98 100644 --- a/tests/unit_tests/distributed/mfsdp_v2/test_cuda_graph.py +++ b/tests/unit_tests/distributed/mfsdp_v2/test_cuda_graph.py @@ -4,7 +4,6 @@ import logging -import pytest import torch from torch import nn from torch.distributed.device_mesh import init_device_mesh @@ -18,6 +17,22 @@ logger = logging.getLogger(__name__) +class NestedModel(nn.Module): + """Model with a root FSDP unit and multiple child FSDP units.""" + + def __init__(self, dim: int, num_children: int) -> None: + super().__init__() + self.bias = nn.Parameter(torch.zeros(dim)) + self.layers = nn.ModuleList([nn.Linear(dim, dim, bias=False) for _ in range(num_children)]) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Run through every child layer with a root-owned bias.""" + x = x + self.bias + for layer in self.layers: + x = layer(x) + return x + + def _flat_placements() -> Placements: return Placements(dp_axes=[0], parameter=[Flat()], gradient=[Flat()], optimizer=[Flat()]) @@ -26,20 +41,21 @@ def test_captures_full_iteration(distributed_setup): """A full training iteration should be CUDA-graphable.""" world_size = distributed_setup.world_size device = distributed_setup.device - if world_size < 2: - pytest.skip("This test requires at least 2 ranks.") mesh = init_device_mesh(device.type, (world_size,)) torch.manual_seed(1234) - model = nn.Linear(4, 2, bias=False).to(device) + dim = 8 + model = NestedModel(dim=dim, num_children=2).to(device) - fully_shard(model, mesh=mesh, placements=_flat_placements()) - optimizer = torch.optim.SGD(model.parameters(), lr=0.25, foreach=False) + static_input = torch.eye(dim, device=device) + static_target = torch.zeros_like(static_input) - static_input = torch.eye(4, device=device) - static_target = torch.tensor( - [[1.0, -0.5], [-0.25, 0.75], [0.5, 0.25], [-0.75, -1.0]], device=device - ) + placements = _flat_placements() + for layer in model.layers: + fully_shard(layer, mesh=mesh, placements=placements) + fully_shard(model, mesh=mesh, placements=placements) + + optimizer = torch.optim.SGD(model.parameters(), lr=0.25, foreach=False) def train_iteration() -> torch.Tensor: optimizer.zero_grad(set_to_none=False) @@ -49,22 +65,26 @@ def train_iteration() -> torch.Tensor: optimizer.step() return loss.detach() - warmup_stream = torch.cuda.Stream() - warmup_stream.wait_stream(torch.cuda.current_stream()) - # Warm up before capture. torch.cuda.graph() uses an internal side stream - # when `stream` is omitted, so `stream=` is only needed when callers must - # control the capture stream, such as when reusing an explicit stream with - # a shared graph memory pool across captures. - with torch.cuda.stream(warmup_stream): + capture_stream = torch.cuda.Stream() + capture_stream.wait_stream(torch.cuda.current_stream()) + + # Warmup + with torch.cuda.stream(capture_stream): + # See: https://docs.nvidia.com/dl-cuda-graph/troubleshooting/memory-issues.html#gradient-accumulator-cross-stream-memory-growth + # Warm up on the same stream used for capture so autograd's accumulation + # path does not create cross-stream gradient-memory growth. # The first warmup installs the reusable sharded gradient views; subsequent # iterations zero them in place for CUDA graph replay. for _ in range(3): train_iteration() + # Capture graph = torch.cuda.CUDAGraph() - with torch.cuda.graph(graph): + with torch.cuda.graph(graph, stream=capture_stream): static_loss = train_iteration() + torch.cuda.current_stream().wait_stream(capture_stream) + # Replay losses = [] for _ in range(5): graph.replay() diff --git a/tests/unit_tests/distributed/mfsdp_v2/test_dbuffer.py b/tests/unit_tests/distributed/mfsdp_v2/test_dbuffer.py index 8631113d480..1180cd55c80 100644 --- a/tests/unit_tests/distributed/mfsdp_v2/test_dbuffer.py +++ b/tests/unit_tests/distributed/mfsdp_v2/test_dbuffer.py @@ -178,6 +178,29 @@ def test_cast_preserves_layout_and_casts_values(distributed_setup): ) +def test_cast_with_out_reuses_destination_and_casts_values(distributed_setup): + """DBuffer.cast writes casted values into an existing destination buffer.""" + mesh = init_device_mesh(distributed_setup.device.type, (distributed_setup.world_size,)) + tensors = _same_tensors_on_all_ranks(distributed_setup.device) + buffer = DBuffer.distribute_tensors(tensors, mesh, [Replicate()]) + destination = DBuffer( + mesh=mesh, + placements=[Replicate()], + tensor_shapes=buffer.layout.tensor_shapes, + dtype=torch.bfloat16, + device=distributed_setup.device, + ) + destination_data_ptr = destination.local_buffer.data_ptr() + + result = buffer.cast(torch.bfloat16, out=destination) + + assert result is destination + assert destination.local_buffer.data_ptr() == destination_data_ptr + _assert_dbuffer_local_tensors_close( + destination, [tensor.to(dtype=torch.bfloat16) for tensor in tensors] + ) + + def test_release_and_reallocate_storage_preserves_buffer_views(distributed_setup): """DBuffer storage can be released and reallocated without replacing existing views.""" mesh = init_device_mesh(distributed_setup.device.type, (distributed_setup.world_size,)) diff --git a/tests/unit_tests/distributed/mfsdp_v2/test_fully_shard.py b/tests/unit_tests/distributed/mfsdp_v2/test_fully_shard.py index 229cd0bff4b..510690d09ec 100644 --- a/tests/unit_tests/distributed/mfsdp_v2/test_fully_shard.py +++ b/tests/unit_tests/distributed/mfsdp_v2/test_fully_shard.py @@ -6,24 +6,32 @@ import pytest import torch +import torch.distributed as dist from torch import nn -from torch.distributed.device_mesh import init_device_mesh +from torch.distributed.device_mesh import DeviceMesh, init_device_mesh from torch.distributed.tensor import DTensor from torch.profiler import ProfilerActivity, profile from megatron.core.distributed.fsdp.src.megatron_fsdp.experimental import ( Flat, + Partial, Placements, + Replicate, fully_shard, + fully_shard_optimizer, microbatch, ) from megatron.core.distributed.fsdp.src.megatron_fsdp.mixed_precision import MixedPrecisionPolicy +from tests.unit_tests.distributed.mfsdp_v2.profiler_utils import ( + collect_linked_kernels, + events_overlap, +) logger = logging.getLogger(__name__) class TinyModel(nn.Module): - """Small model with two separately shardable units.""" + """Small model with two separately shardable modules.""" def __init__(self) -> None: super().__init__() @@ -50,7 +58,7 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: class MultiChildModel(nn.Module): - """Model with direct parameters and multiple child FSDP units.""" + """Model with direct parameters and multiple child FsdpModules.""" def __init__(self, dim: int, num_children: int) -> None: super().__init__() @@ -100,19 +108,32 @@ def _flat_placements() -> Placements: return Placements(dp_axes=[0], parameter=[Flat()], gradient=[Flat()], optimizer=[Flat()]) +def _hsdp_placements() -> Placements: + """HSDP: params/optimizer replicated across DP-outer (axis 0), sharded within + DP-inner (axis 1). main_grad rests [Partial, Flat] between microbatches and is + all-reduced to [Replicate, Flat] on the last microbatch.""" + return Placements( + dp_axes=[0, 1], + parameter=[Replicate(), Flat()], + gradient=[Partial(dist.ReduceOp.AVG), Flat()], + optimizer=[Replicate(), Flat()], + ) + + def _mb(num_bytes: int) -> str: return f"{num_bytes / 1024**2:.2f} MB" -def _events_overlap(first, second) -> bool: - return ( - first.time_range.start < second.time_range.end - and second.time_range.start < first.time_range.end - ) +# CPU ops that a device event chains up to via cpu_parent, used to attribute the device +# work to its enclosing collective or matmul operation. +_ALL_GATHER_OP_NAME_SUBSTRING = "allgather" +_REDUCE_SCATTER_OP_NAME_SUBSTRING = "reduce_scatter" +_ALLREDUCE_OP_NAME_SUBSTRING = "allreduce" +_GEMM_OP_NAME_SUBSTRING = "aten::mm" @pytest.mark.parametrize("num_microbatches", [1, 3]) -def test_fully_shard_losses_match_baseline(distributed_setup, num_microbatches): +def test_fully_shard_sgd_losses_match_baseline(distributed_setup, num_microbatches): """Minimal per-module FSDP training should match single-rank SGD.""" rank = distributed_setup.rank world_size = distributed_setup.world_size @@ -168,8 +189,148 @@ def train(model, optimizer, log_prefix) -> list[torch.Tensor]: ) +@pytest.mark.parametrize("set_to_none", [True, False]) +@pytest.mark.parametrize("num_microbatches", [1, 3]) +def test_hsdp_losses_match_baseline(distributed_setup, num_microbatches, set_to_none): + """HSDP (DP-outer replicated, DP-inner sharded) training should match single-rank SGD. + + Gradients reduce-scatter within DP-inner every backward and accumulate into + main_grad; the DP-outer all-reduce runs only on the last microbatch, scoped + via ``microbatch(...)``. Every rank sees identical data, so the averaged + gradient equals the single-rank gradient and losses must match. Both + ``zero_grad`` modes are covered: ``set_to_none=True`` overwrites main_grad, + ``set_to_none=False`` accumulates into a zeroed main_grad. + """ + rank = distributed_setup.rank + world_size = distributed_setup.world_size + device = distributed_setup.device + if world_size < 4 or world_size % 2 != 0: + pytest.skip("This test requires an even number of at least 4 ranks for a 2-D DP mesh.") + + outer_size = 2 + inner_size = world_size // outer_size + mesh = init_device_mesh( + device.type, (outer_size, inner_size), mesh_dim_names=("dp_outer", "dp_inner") + ) + torch.manual_seed(1234) + dim = 8 + baseline = MultiChildModel(dim=dim, num_children=2).to(device) + model = MultiChildModel(dim=dim, num_children=2).to(device) + model.load_state_dict(baseline.state_dict()) + + # Shard the child layers, then the model, so the children share a root context + # and reduce through the overlap path instead of as independent roots. + for layer in model.layers: + fully_shard(layer, mesh=mesh, placements=_hsdp_placements()) + fully_shard(model, mesh=mesh, placements=_hsdp_placements()) + baseline_optimizer = torch.optim.SGD(baseline.parameters(), lr=0.05) + optimizer = torch.optim.SGD(model.parameters(), lr=0.05) + + micro_batch_size = 2 + x = torch.randn(num_microbatches, micro_batch_size, dim, device=device) + target = torch.randn(num_microbatches, micro_batch_size, dim, device=device) + microbatches = tuple(zip(x.unbind(), target.unbind())) + + def train(model, optimizer, log_prefix) -> list[torch.Tensor]: + losses = [] + for step in range(5): + optimizer.zero_grad(set_to_none=set_to_none) + + for microbatch_index, (microbatch_x, microbatch_target) in enumerate(microbatches): + is_last = microbatch_index == num_microbatches - 1 + with microbatch(model, is_last=is_last): + loss = torch.nn.functional.mse_loss(model(microbatch_x), microbatch_target) + (loss / num_microbatches).backward() + losses.append(loss.detach()) + logger.debug( + "%s train parity: rank=%s, step=%s, microbatch=%s, loss=%s", + log_prefix, + rank, + step, + microbatch_index, + loss, + ) + + optimizer.step() + return losses + + baseline_losses = train(baseline, baseline_optimizer, "Baseline") + sharded_losses = train(model, optimizer, "HSDP") + + torch.testing.assert_close( + torch.stack(sharded_losses), + torch.stack(baseline_losses), + msg="HSDP losses did not match baseline losses.", + ) + + +def test_hsdp_defers_dp_outer_allreduce_to_last_microbatch(distributed_setup): + """HSDP reduce-scatters DP-inner every microbatch but all-reduces DP-outer once. + + ``fully_shard(model)`` makes the child units share a root context so their + reductions run through the overlap path rather than as independent roots. + Counting linked NCCL kernels over a multi-microbatch step, the DP-inner reduce-scatter + fires once per microbatch per group while the DP-outer all-reduce that + finalizes main_grad fires only on the last microbatch, so the reduce-scatter + count is exactly ``num_microbatches`` times the all-reduce count. This asserts + on kernel counts only, not numerics. + """ + world_size = distributed_setup.world_size + device = distributed_setup.device + if world_size < 4 or world_size % 2 != 0: + pytest.skip("This test requires an even number of at least 4 ranks for a 2-D DP mesh.") + + outer_size = 2 + inner_size = world_size // outer_size + mesh = init_device_mesh( + device.type, (outer_size, inner_size), mesh_dim_names=("dp_outer", "dp_inner") + ) + torch.manual_seed(1234) + dim = 8 + num_children = 2 + model = MultiChildModel(dim=dim, num_children=num_children).to(device) + for layer in model.layers: + fully_shard(layer, mesh=mesh, placements=_hsdp_placements()) + fully_shard(model, mesh=mesh, placements=_hsdp_placements()) + optimizer = torch.optim.SGD(model.parameters(), lr=0.05) + + num_microbatches = 3 + micro_batch_size = 2 + x = torch.randn(num_microbatches, micro_batch_size, dim, device=device) + target = torch.randn(num_microbatches, micro_batch_size, dim, device=device) + microbatches = tuple(zip(x.unbind(), target.unbind())) + + def train_one_step() -> None: + optimizer.zero_grad(set_to_none=True) + for microbatch_index, (microbatch_x, microbatch_target) in enumerate(microbatches): + is_last = microbatch_index == num_microbatches - 1 + with microbatch(model, is_last=is_last): + loss = torch.nn.functional.mse_loss(model(microbatch_x), microbatch_target) + (loss / num_microbatches).backward() + optimizer.step() + + train_one_step() + torch.cuda.synchronize(device) + + with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof: + train_one_step() + torch.cuda.synchronize(device) + + reduce_scatter_kernels = collect_linked_kernels(prof, _REDUCE_SCATTER_OP_NAME_SUBSTRING) + allreduce_kernels = collect_linked_kernels(prof, _ALLREDUCE_OP_NAME_SUBSTRING) + # One DP-outer all-reduce per parameter group -- each child layer plus the + # root unit's bias -- fired only on the last microbatch. Plain DP fires none. + assert len(allreduce_kernels) == num_children + 1, [event.name for event in prof.events()] + # DP-inner reduce-scatter runs every microbatch; the DP-outer all-reduce runs + # only on the last, so the counts differ by exactly the microbatch factor. + assert len(reduce_scatter_kernels) == len(allreduce_kernels) * num_microbatches, ( + f"Expected reduce-scatter ({len(reduce_scatter_kernels)}) to be {num_microbatches}x " + f"the DP-outer all-reduce count ({len(allreduce_kernels)})." + ) + + def test_nested_fully_shard_excludes_child_owned_parameters(distributed_setup): - """An outer FSDP unit owns direct parameters but not nested child-unit parameters.""" + """An outer FsdpModule owns direct parameters but not nested child FsdpModule parameters.""" world_size = distributed_setup.world_size device = distributed_setup.device if world_size < 2: @@ -311,23 +472,69 @@ def test_root_backward_returns_to_resting_memory(distributed_setup): ) -def test_overlaps_all_gather_and_compute(distributed_setup): - """A shared root context should let child all-gathers overlap GEMM compute.""" +@pytest.mark.parametrize("use_symm_mem", [False, True], ids=["default", "symmetric_memory"]) +def test_overlaps_communication_and_compute(distributed_setup, use_symm_mem): + """Forward and backward communication should overlap GEMM compute.""" world_size = distributed_setup.world_size device = distributed_setup.device if world_size < 2: pytest.skip("This test requires at least 2 ranks.") - mesh = init_device_mesh(device.type, (world_size,)) - dim = 4096 + # A large hidden size keeps the per-layer GEMMs long enough that the + # collectives reliably overlap them. The overlap count is otherwise + # launch-bound: the host issues kernels with gaps (amplified by CI's + # `coverage run` wrapper), so with short GEMMs a collective can land in a + # gap between GEMMs instead of running alongside one, making the count jitter + # run to run. At dim=16384 the GEMMs dominate that launch jitter and the + # overlap becomes deterministic. (dim=8192 was flaky under coverage.) + dim = 16384 num_children = 4 dtype = torch.bfloat16 + + # new_group requires a default process group. Initialize it here so this test works + # in isolation. Do not eagerly initialize it with device_id in the shared fixture: + # that can hang teardown after communicator splits; see + # https://github.com/pytorch/pytorch/issues/190396. + if not dist.is_initialized(): + dist.init_process_group(backend="nccl") + + if use_symm_mem: + # Dedicated communicator with NCCL's zero-CTA policy. cta_policy is a + # per-communicator property, so scoping it to this group leaves the rest of the + # bucket on default-CTA symmetric-memory kernels (test_symmetric_memory.py asserts + # ncclSymk all-gather kernel counts, which zero-CTA would turn into copy-engine + # memcpys). This 1-D group models the DP (FSDP) sub-mesh that mfsdp is handed in + # production: with EP/TP the full device mesh is multi-dimensional, but mfsdp + # requires an all-FSDP mesh (see experimental/module.py) and never sees the TP/EP + # axes, so only the DP communicator needs the zero-CTA policy. + zero_cta_options = dist.ProcessGroupNCCL.Options() + zero_cta_options.config.cta_policy = dist.ProcessGroupNCCL.NCCL_CTA_POLICY_ZERO + dp_group = dist.new_group(backend="nccl", pg_options=zero_cta_options) + # NCCL window registration can fail when symmetric-memory rendezvous is the first + # operation on a communicator, so initialize this communicator explicitly. + dist.barrier(group=dp_group, device_ids=[device.index]) + else: + dp_group = dist.new_group(backend="nccl") + + mesh = DeviceMesh.from_group(dp_group, device.type) model = MultiChildModel(dim=dim, num_children=num_children).to(dtype=dtype) placements = _flat_placements() policy = MixedPrecisionPolicy(main_params_dtype=dtype, main_grads_dtype=dtype) for layer in model.layers: - fully_shard(layer, mesh=mesh, placements=placements, mixed_precision_policy=policy) - fully_shard(model, mesh=mesh, placements=placements, mixed_precision_policy=policy) + fully_shard( + layer, + mesh=mesh, + placements=placements, + mixed_precision_policy=policy, + use_symm_mem=use_symm_mem, + ) + fully_shard( + model, + mesh=mesh, + placements=placements, + mixed_precision_policy=policy, + use_symm_mem=use_symm_mem, + ) x = torch.randn(4096, dim, device=device, dtype=dtype, requires_grad=True) @@ -346,45 +553,87 @@ def train_one_iteration() -> None: # drop the CUDA events. torch.cuda.synchronize(device) - cuda_events = [event for event in prof.events() if event.device_type.name == "CUDA"] - all_gather_events = [ - event - for event in cuda_events - if "nccl" in event.name.lower() and "allgather" in event.name.lower() - ] - # GEMM device-kernel names vary across CUDA/cuBLAS versions and GPU archs - # (e.g. "*gemm*", "cutlass*", "cublas*", and cuBLASLt's Hopper "nvjet_sm90_*"). - gemm_events = [ - event - for event in cuda_events - if any(token in event.name.lower() for token in ("gemm", "cutlass", "cublas", "nvjet")) - ] - assert all_gather_events, [event.name for event in cuda_events] - assert gemm_events, [event.name for event in cuda_events] - - all_gather_streams = {event.device_resource_id for event in all_gather_events} - gemm_streams = {event.device_resource_id for event in gemm_events} - assert len(all_gather_streams) == 1 - assert all_gather_streams.isdisjoint(gemm_streams) - - overlap_count = sum( - any(_events_overlap(all_gather_event, gemm_event) for gemm_event in gemm_events) - for all_gather_event in all_gather_events + gemm_kernels = collect_linked_kernels(prof, _GEMM_OP_NAME_SUBSTRING) + # Each child Linear runs one forward and two backward matmuls. aten::mm may also + # launch auxiliary kernels, so check only the matmul lower bound. + assert len(gemm_kernels) >= 3 * num_children, ( + f"Expected at least {3 * num_children} kernels linked to GEMMs, got " + f"{len(gemm_kernels)}: " + f"{[kernel.name for kernel in gemm_kernels]}" ) - # This profiles a full forward/backward iteration, so backward all-gathers are - # included in all_gather_events. The expected overlap count is from the forward - # child pipeline: each child after the first can all-gather while the previous - # child computes, giving num_children - 1 overlaps. Backward does not overlap - # in this all-gather-only path because gradient reduction is not delayed: - # each module synchronously reduces gradients in post_backward before autograd - # reaches the next module's pre_backward all-gather. The next PR addresses - # this by delaying gradient reduction. - expected_overlap_count = num_children - 1 - assert overlap_count >= expected_overlap_count, ( - f"Expected at least {expected_overlap_count} all-gather events to overlap compute, " - f"got {overlap_count}/{len(all_gather_events)}." + + allgather_kernels = collect_linked_kernels(prof, _ALL_GATHER_OP_NAME_SUBSTRING) + reduce_scatter_kernels = collect_linked_kernels(prof, _REDUCE_SCATTER_OP_NAME_SUBSTRING) + # The num_children child layers plus the root are each a sharded module; each does a + # forward and a backward all-gather and one reduce-scatter. Zero-CTA moves the + # all-gather to copy-engine memcpys, so it should not emit all-gather kernels. + num_sharded_modules = num_children + 1 + expected_allgather_kernel_count = 0 if use_symm_mem else 2 * num_sharded_modules + assert len(allgather_kernels) == expected_allgather_kernel_count, ( + f"Expected {expected_allgather_kernel_count} all-gather kernels, got " + f"{len(allgather_kernels)}: {[kernel.name for kernel in allgather_kernels]}" + ) + assert len(reduce_scatter_kernels) == num_sharded_modules, ( + f"Expected {num_sharded_modules} reduce-scatter kernels, got " + f"{len(reduce_scatter_kernels)}: {[kernel.name for kernel in reduce_scatter_kernels]}" ) + allgather_streams = {kernel.device_resource_id for kernel in allgather_kernels} + reduce_scatter_streams = {kernel.device_resource_id for kernel in reduce_scatter_kernels} + gemm_streams = {kernel.device_resource_id for kernel in gemm_kernels} + if allgather_kernels: + assert len(allgather_streams) == 1 + assert len(reduce_scatter_streams) == 1 + assert allgather_streams.isdisjoint(reduce_scatter_streams) + assert allgather_streams.isdisjoint(gemm_streams) + assert reduce_scatter_streams.isdisjoint(gemm_streams) + + allgather_overlap_count = sum( + any(events_overlap(kernel, gemm) for gemm in gemm_kernels) for kernel in allgather_kernels + ) + reduce_scatter_overlap_count = sum( + any(events_overlap(kernel, gemm) for gemm in gemm_kernels) + for kernel in reduce_scatter_kernels + ) + expected_allgather_overlap = 2 * (num_children - 1) + expected_reduce_scatter_overlap = num_children - 1 + if not use_symm_mem: + assert allgather_overlap_count >= expected_allgather_overlap, ( + f"Expected at least {expected_allgather_overlap} all-gathers to " + f"overlap compute, got {allgather_overlap_count}/{len(allgather_kernels)}." + ) + assert reduce_scatter_overlap_count >= expected_reduce_scatter_overlap, ( + f"Expected at least {expected_reduce_scatter_overlap} reduce-scatters to overlap " + f"compute, got {reduce_scatter_overlap_count}/{len(reduce_scatter_kernels)}." + ) + + # Release the dedicated communicator so it does not leak into the shared session. + dist.destroy_process_group(dp_group) + + +def test_parameterless_parent_with_child_modules_trains(distributed_setup): + """A parent with no unowned parameters should still root trainable child FsdpModules.""" + world_size = distributed_setup.world_size + device = distributed_setup.device + + mesh = init_device_mesh(device.type, (world_size,)) + torch.manual_seed(5678) + model = nn.Sequential(nn.Linear(4, 4, bias=False), nn.Linear(4, 2, bias=False)).to(device) + + fully_shard(model[0], mesh=mesh, placements=_flat_placements()) + fully_shard(model[1], mesh=mesh, placements=_flat_placements()) + fully_shard(model, mesh=mesh, placements=_flat_placements()) + + assert model.parameter_groups == () + + optimizer = torch.optim.SGD(model.parameters(), lr=0.05) + x = torch.randn(3, 4, device=device) + + optimizer.zero_grad(set_to_none=True) + loss = model(x).sum() + loss.backward() + optimizer.step() + def test_frozen_parameter_group_does_not_allocate_main_grad(distributed_setup): """A non-trainable parameter group should not allocate persistent main gradients.""" @@ -450,6 +699,7 @@ def test_next_forward_uses_optimizer_updated_weights(distributed_setup): # SGD's foreach/fused CUDA paths require matching parameter and gradient dtypes. # Use the scalar path to exercise FP32 main weights with default BF16 main grads. optimizer = torch.optim.SGD(model.parameters(), lr=0.25, foreach=False) + fully_shard_optimizer(optimizer) x = torch.ones(1, 1, device=device, dtype=torch.bfloat16) def train_iteration() -> torch.Tensor: @@ -466,6 +716,83 @@ def train_iteration() -> torch.Tensor: torch.testing.assert_close(second_loss, first_loss) +def test_optimizer_post_step_syncs_once_per_parameter_group(distributed_setup, monkeypatch): + """Optimizer synchronization should run once per group, not once per microbatch.""" + world_size = distributed_setup.world_size + device = distributed_setup.device + if world_size < 2: + pytest.skip("This test requires at least 2 ranks.") + + mesh = init_device_mesh(device.type, (world_size,)) + model = TinyModel().to(device=device, dtype=torch.bfloat16) + fully_shard(model.fc1, mesh=mesh, placements=_flat_placements()) + fully_shard(model.fc2, mesh=mesh, placements=_flat_placements()) + parameter_groups = (*model.fc1.parameter_groups, *model.fc2.parameter_groups) + sync_counts = {parameter_group: 0 for parameter_group in parameter_groups} + + def make_count_sync(parameter_group): + sync_model_weight = parameter_group.sync_model_weight_from_main_weight + + def count_sync(): + sync_counts[parameter_group] += 1 + sync_model_weight() + + return count_sync + + for parameter_group in parameter_groups: + monkeypatch.setattr( + parameter_group, "sync_model_weight_from_main_weight", make_count_sync(parameter_group) + ) + + optimizer = torch.optim.Adam(model.parameters(), lr=0.01) + fully_shard_optimizer(optimizer) + inputs = torch.randn(3, 2, 8, device=device, dtype=torch.bfloat16) + + for step in range(3): + optimizer.zero_grad(set_to_none=True) + for microbatch_input in inputs: + (model(microbatch_input).sum() / len(inputs)).backward() + + assert all(sync_count == step for sync_count in sync_counts.values()) + optimizer.step() + assert all(sync_count == step + 1 for sync_count in sync_counts.values()) + + +def test_fully_shard_adam_mixed_precision_losses_match_baseline(distributed_setup): + """Mixed-precision FSDP Adam should track an unsharded Adam baseline.""" + world_size = distributed_setup.world_size + device = distributed_setup.device + if world_size < 2: + pytest.skip("This test requires at least 2 ranks.") + mesh = init_device_mesh(device.type, (world_size,)) + torch.manual_seed(2026) + baseline = TinyModel().to(device=device, dtype=torch.bfloat16) + model = TinyModel().to(device=device, dtype=torch.bfloat16) + model.load_state_dict(baseline.state_dict()) + fully_shard(model.fc1, mesh=mesh, placements=_flat_placements()) + fully_shard(model.fc2, mesh=mesh, placements=_flat_placements()) + + baseline_optimizer = torch.optim.Adam(baseline.parameters(), lr=0.01) + optimizer = torch.optim.Adam(model.parameters(), lr=0.01) + fully_shard_optimizer(optimizer) + + x = torch.randn(3, 8, device=device, dtype=torch.bfloat16) + target = torch.randn(3, 4, device=device, dtype=torch.bfloat16) + + for _ in range(3): + baseline_optimizer.zero_grad() + optimizer.zero_grad() + + baseline_loss = torch.nn.functional.mse_loss(baseline(x).float(), target.float()) + loss = torch.nn.functional.mse_loss(model(x).float(), target.float()) + torch.testing.assert_close(loss, baseline_loss, rtol=0, atol=3e-3) + + baseline_loss.backward() + loss.backward() + baseline_optimizer.step() + optimizer.step() + + def test_microbatch_scopes_child_contexts(distributed_setup): """microbatch() should scope FSDP child contexts under an unwrapped parent.""" world_size = distributed_setup.world_size @@ -485,24 +812,30 @@ def test_microbatch_scopes_child_contexts(distributed_setup): def test_cpu_initialized_parameters_shard_to_mesh_device(distributed_setup): - """CPU-initialized parameters should be sharded with their real values.""" + """A CPU model should support sharding a child before moving the full model to CUDA.""" world_size = distributed_setup.world_size device = distributed_setup.device - if world_size < 2: - pytest.skip("This test requires at least 2 ranks.") mesh = init_device_mesh(device.type, (world_size,)) - model = nn.Linear(4, 4, bias=False) + model = nn.Sequential(nn.Linear(4, 4, bias=False), nn.Linear(4, 4, bias=False)) with torch.no_grad(): - model.weight.fill_(3.0) - expected_weight = model.weight.detach().to(device) + model[0].weight.fill_(2.0) + model[1].weight.fill_(3.0) + x = torch.ones(1, 4) + expected_output = model(x).to(device) - fully_shard(model, mesh=mesh, placements=_flat_placements()) + # Shard the second layer's parameters onto the mesh device; the unwrapped + # first layer's parameters remain on CPU until model.to(device) below. + fully_shard(model[1], mesh=mesh, placements=_flat_placements()) - (group,) = model.parameter_groups - full_weight = group.model_weight.allgather(0).get_local_tensor(0) - assert full_weight.device.type == device.type - torch.testing.assert_close(full_weight, expected_weight) + assert model[0].weight.device.type == "cpu" + assert isinstance(model[1].weight, DTensor) + assert model[1].weight.device == device + + model.to(device) + + output = model(x.to(device)) + torch.testing.assert_close(output, expected_output) def test_non_leaf_parameter_view_survives_storage_resize(distributed_setup): diff --git a/tests/unit_tests/distributed/mfsdp_v2/test_optimizer.py b/tests/unit_tests/distributed/mfsdp_v2/test_optimizer.py new file mode 100644 index 00000000000..4ba4d1be271 --- /dev/null +++ b/tests/unit_tests/distributed/mfsdp_v2/test_optimizer.py @@ -0,0 +1,101 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Unit tests for Megatron-FSDP optimizer behavior.""" + +import pytest +import torch +from torch import nn +from torch.distributed.device_mesh import init_device_mesh +from transformer_engine.pytorch.optimizers import FusedAdam + +from megatron.core.distributed.fsdp.src.megatron_fsdp.experimental import ( + Flat, + Placements, + fully_shard, + fully_shard_optimizer, +) +from megatron.core.distributed.fsdp.src.megatron_fsdp.mixed_precision import MixedPrecisionPolicy + + +class TinyModel(nn.Module): + """Small model with two separately shardable units.""" + + def __init__(self) -> None: + super().__init__() + self.fc1 = nn.Linear(8, 16) + self.relu = nn.ReLU() + self.fc2 = nn.Linear(16, 4) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + """Run the tiny model.""" + return self.fc2(self.relu(self.fc1(x))) + + +def _flat_placements() -> Placements: + return Placements(dp_axes=[0], parameter=[Flat()], gradient=[Flat()], optimizer=[Flat()]) + + +def test_adam_without_adapter_raises_precision_error(distributed_setup): + """Raw Adam should fail on mixed-precision FSDP parameters without the adapter.""" + world_size = distributed_setup.world_size + device = distributed_setup.device + mesh = init_device_mesh(device.type, (world_size,)) + torch.manual_seed(2026) + model = TinyModel().to(device=device, dtype=torch.bfloat16) + fully_shard(model.fc1, mesh=mesh, placements=_flat_placements()) + fully_shard(model.fc2, mesh=mesh, placements=_flat_placements()) + optimizer = torch.optim.Adam(model.parameters(), lr=0.01) + + x = torch.randn(6, 8, device=device, dtype=torch.bfloat16) + optimizer.zero_grad(set_to_none=True) + loss = model(x).sum() + loss.backward() + + with pytest.raises(RuntimeError, match="dtype"): + optimizer.step() + + +def test_fused_adam_adapter_accepts_mismatched_grads(distributed_setup): + """TE FusedAdam should handle mixed-precision FSDP grads through the adapter.""" + world_size = distributed_setup.world_size + device = distributed_setup.device + + mesh = init_device_mesh(device.type, (world_size,)) + torch.manual_seed(2026) + model = TinyModel().to(device=device, dtype=torch.bfloat16) + # These are the defaults, but spell them out so the test clearly exercises + # mismatched parameter and gradient precision. + mixed_precision_policy = MixedPrecisionPolicy( + main_params_dtype=torch.float32, main_grads_dtype=torch.bfloat16 + ) + fully_shard( + model.fc1, + mesh=mesh, + placements=_flat_placements(), + mixed_precision_policy=mixed_precision_policy, + ) + fully_shard( + model.fc2, + mesh=mesh, + placements=_flat_placements(), + mixed_precision_policy=mixed_precision_policy, + ) + optimizer = FusedAdam(model.parameters(), lr=0.01) + fully_shard_optimizer(optimizer, precision_aware=True) + + x = torch.randn(6, 8, device=device, dtype=torch.bfloat16) + optimizer.zero_grad(set_to_none=True) + loss = model(x).sum() + loss.backward() + + for parameter in model.parameters(): + assert parameter.grad is not None + assert parameter.dtype != parameter.grad.dtype + + params_before_step = [parameter.detach().clone() for parameter in model.parameters()] + optimizer.step() + + assert any( + not torch.equal(parameter_before, parameter.detach()) + for parameter_before, parameter in zip(params_before_step, model.parameters()) + ) diff --git a/tests/unit_tests/distributed/mfsdp_v2/test_symmetric_memory.py b/tests/unit_tests/distributed/mfsdp_v2/test_symmetric_memory.py index 5e84ad77c75..87207548140 100644 --- a/tests/unit_tests/distributed/mfsdp_v2/test_symmetric_memory.py +++ b/tests/unit_tests/distributed/mfsdp_v2/test_symmetric_memory.py @@ -4,8 +4,9 @@ import pytest import torch +import torch.distributed as dist from torch import nn -from torch.distributed.device_mesh import init_device_mesh +from torch.distributed.device_mesh import DeviceMesh, init_device_mesh from torch.profiler import ProfilerActivity, profile from megatron.core.distributed.fsdp.src.megatron_fsdp import MixedPrecisionPolicy @@ -14,17 +15,20 @@ Placements, fully_shard, ) +from tests.unit_tests.distributed.mfsdp_v2.profiler_utils import collect_linked_kernels # Each sharded Linear's collective must be large enough that NCCL selects its # symmetric-memory (ncclSymk*) kernels over ring. Sub-KB collectives fall back to # ring on some platforms (e.g. CI with NCCL_NVLS_ENABLE=0), which would make the -# symmetric-kernel assertions below fail; 1024-wide units (a few-MiB bf16 weight) +# symmetric-kernel assertions below fail; 1024-wide layers (a few-MiB bf16 weight) # reliably engage the symmetric kernels. _HIDDEN = 1024 +_ALL_GATHER_OP_NAME_SUBSTRING = "allgather" +_REDUCE_SCATTER_OP_NAME_SUBSTRING = "reduce_scatter" class TinyModel(nn.Module): - """Two separately shardable units, sized so NCCL selects symmetric-memory kernels.""" + """Two separately shardable Linear modules, sized so NCCL selects symmetric-memory kernels.""" def __init__(self) -> None: super().__init__() @@ -41,18 +45,6 @@ def _flat_placements() -> Placements: return Placements(dp_axes=[0], parameter=[Flat()], gradient=[Flat()], optimizer=[Flat()]) -def _kernels(prof: torch.profiler.profile) -> list[str]: - return [event.name for event in prof.events()] - - -def _is_symmetric_kernel(kernel: str) -> bool: - return "ncclSymk" in kernel - - -def _count_symmetric_kernels(kernels: list[str], subname: str) -> int: - return sum(1 for kernel in kernels if _is_symmetric_kernel(kernel) and subname in kernel) - - @pytest.mark.parametrize("num_microbatches", [1, 3]) def test_fully_shard_symmetric_memory_matches_default_and_profiles_nccl( distributed_setup, num_microbatches @@ -106,11 +98,11 @@ def train(use_symm_mem: bool) -> list[torch.Tensor]: return losses - with profile(activities=[ProfilerActivity.CUDA]) as prof_without_symm_mem: + with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof_without_symm_mem: losses_without_symm_mem = train(use_symm_mem=False) torch.cuda.synchronize() - with profile(activities=[ProfilerActivity.CUDA]) as prof_with_symm_mem: + with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof_with_symm_mem: losses_with_symm_mem = train(use_symm_mem=True) torch.cuda.synchronize() @@ -120,29 +112,125 @@ def train(use_symm_mem: bool) -> list[torch.Tensor]: msg="Symmetric-memory FSDP losses did not match default FSDP losses.", ) - kernels_without_symm_mem = _kernels(prof_without_symm_mem) - assert _count_symmetric_kernels(kernels_without_symm_mem, "AllGather") == 0 - assert _count_symmetric_kernels(kernels_without_symm_mem, "ReduceScatter") == 0 + allgather_kernels_without_symm_mem = collect_linked_kernels( + prof_without_symm_mem, _ALL_GATHER_OP_NAME_SUBSTRING + ) + reduce_scatter_kernels_without_symm_mem = collect_linked_kernels( + prof_without_symm_mem, _REDUCE_SCATTER_OP_NAME_SUBSTRING + ) + assert all("ncclSymk" not in kernel.name for kernel in allgather_kernels_without_symm_mem) + assert all("ncclSymk" not in kernel.name for kernel in reduce_scatter_kernels_without_symm_mem) - kernels_with_symm_mem = _kernels(prof_with_symm_mem) + allgather_kernels_with_symm_mem = collect_linked_kernels( + prof_with_symm_mem, _ALL_GATHER_OP_NAME_SUBSTRING + ) + reduce_scatter_kernels_with_symm_mem = collect_linked_kernels( + prof_with_symm_mem, _REDUCE_SCATTER_OP_NAME_SUBSTRING + ) # 2 sharded modules (fc1, fc2), one reduce-scatter each per microbatch step. expected_reduce_scatter_kernel_count = num_training_steps * num_microbatches * 2 - nccl_kernels_with_symm_mem = [ - kernel for kernel in kernels_with_symm_mem if "nccl" in kernel.lower() - ] - assert ( - _count_symmetric_kernels(kernels_with_symm_mem, "ReduceScatter") - == expected_reduce_scatter_kernel_count - ), ( + assert len(reduce_scatter_kernels_with_symm_mem) == expected_reduce_scatter_kernel_count, ( "Unexpected NCCL symmetric-memory reduce-scatter kernel count. " - f"Observed NCCL kernels: {nccl_kernels_with_symm_mem[:20]}" + f"Observed reduce-scatter kernels: " + f"{[kernel.name for kernel in reduce_scatter_kernels_with_symm_mem[:20]]}" + ) + assert all("ncclSymk" in kernel.name for kernel in reduce_scatter_kernels_with_symm_mem), ( + "Expected all symmetric-memory reduce-scatter kernels to be ncclSymk kernels. " + f"Observed reduce-scatter kernels: " + f"{[kernel.name for kernel in reduce_scatter_kernels_with_symm_mem[:20]]}" ) - expected_all_gather_kernel_count = 2 * expected_reduce_scatter_kernel_count - assert ( - _count_symmetric_kernels(kernels_with_symm_mem, "AllGather") - == expected_all_gather_kernel_count - ), ( + expected_allgather_kernel_count = 2 * expected_reduce_scatter_kernel_count + assert len(allgather_kernels_with_symm_mem) == expected_allgather_kernel_count, ( "Unexpected NCCL symmetric-memory all-gather kernel count. " - f"Observed NCCL kernels: {nccl_kernels_with_symm_mem[:20]}" + f"Observed all-gather kernels: " + f"{[kernel.name for kernel in allgather_kernels_with_symm_mem[:20]]}" + ) + assert all("ncclSymk" in kernel.name for kernel in allgather_kernels_with_symm_mem), ( + "Expected all symmetric-memory all-gather kernels to be ncclSymk kernels. " + f"Observed all-gather kernels: " + f"{[kernel.name for kernel in allgather_kernels_with_symm_mem[:20]]}" + ) + + +def test_fully_shard_zero_cta_moves_all_gather_to_copy_engine(distributed_setup): + """NCCL's zero-CTA policy runs the all-gather on the copy engine. + + Zero-CTA offloads only pure data movement, so the all-gather emits no ``ncclSymk`` + kernel (it becomes a copy-engine memcpy). The reduce-scatter's reduction cannot run on + the copy engine, so it stays a symmetric-memory kernel -- an SM-launched NVLS multicast + reduce (ncclSymkDevKernel_ReduceScatter_LDMC/LL). + """ + world_size = distributed_setup.world_size + device = distributed_setup.device + if world_size < 2: + pytest.skip("This test requires at least 2 ranks.") + + # new_group requires a default process group. Initialize it here so this test works + # in isolation. Do not eagerly initialize it with device_id in the shared fixture: + # that can hang teardown after communicator splits; see + # https://github.com/pytorch/pytorch/issues/190396. + if not dist.is_initialized(): + dist.init_process_group(backend="nccl") + + # Dedicated communicator with NCCL's zero-CTA policy, scoped to this test so the rest + # of the bucket keeps default-CTA symmetric-memory kernels. + zero_cta_options = dist.ProcessGroupNCCL.Options() + zero_cta_options.config.cta_policy = dist.ProcessGroupNCCL.NCCL_CTA_POLICY_ZERO + dp_group = dist.new_group(backend="nccl", pg_options=zero_cta_options) + # NCCL window registration can fail when symmetric-memory rendezvous is the first + # operation on a communicator, so initialize this communicator explicitly. + dist.barrier(group=dp_group, device_ids=[device.index]) + mesh = DeviceMesh.from_group(dp_group, device.type) + + num_training_steps = 5 + model = TinyModel().to(device=device, dtype=torch.bfloat16) + mixed_precision_policy = MixedPrecisionPolicy(main_params_dtype=torch.float32) + fully_shard( + model.fc1, + mesh=mesh, + placements=_flat_placements(), + mixed_precision_policy=mixed_precision_policy, + use_symm_mem=True, + ) + fully_shard( + model.fc2, + mesh=mesh, + placements=_flat_placements(), + mixed_precision_policy=mixed_precision_policy, + use_symm_mem=True, ) + optimizer = torch.optim.SGD(model.parameters(), lr=0.05, foreach=False) + x = torch.randn(2, _HIDDEN, device=device, dtype=torch.bfloat16) + target = torch.randn(2, _HIDDEN, device=device, dtype=torch.bfloat16) + + with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof: + for _ in range(num_training_steps): + optimizer.zero_grad() + torch.nn.functional.mse_loss(model(x), target).backward() + optimizer.step() + torch.cuda.synchronize() + + allgather_kernels = collect_linked_kernels(prof, _ALL_GATHER_OP_NAME_SUBSTRING) + reduce_scatter_kernels = collect_linked_kernels(prof, _REDUCE_SCATTER_OP_NAME_SUBSTRING) + # Zero-CTA moves the all-gather to the copy engine: no symmetric-memory all-gather kernel. + assert not allgather_kernels, ( + f"Expected no symmetric-memory all-gather kernel under zero-CTA. " + f"Observed all-gather kernels: {[kernel.name for kernel in allgather_kernels[:20]]}" + ) + # The reduce-scatter's reduction cannot run on the copy engine, so it stays a + # symmetric-memory kernel (an SM-launched NVLS multicast reduce): one per sharded + # module (fc1, fc2) per training step. + expected_reduce_scatter_kernel_count = num_training_steps * 2 + assert len(reduce_scatter_kernels) == expected_reduce_scatter_kernel_count, ( + f"Expected {expected_reduce_scatter_kernel_count} symmetric-memory reduce-scatter " + f"kernels under zero-CTA. Observed reduce-scatter kernels: " + f"{[kernel.name for kernel in reduce_scatter_kernels[:20]]}" + ) + assert all("ncclSymk" in kernel.name for kernel in reduce_scatter_kernels), ( + "Expected all zero-CTA reduce-scatter kernels to be ncclSymk kernels. " + f"Observed reduce-scatter kernels: {[kernel.name for kernel in reduce_scatter_kernels[:20]]}" + ) + + # Release the dedicated communicator (leaks only on a test failure above, which is fine). + dist.destroy_process_group(dp_group) diff --git a/tests/unit_tests/distributed/test_finalize_model_grads.py b/tests/unit_tests/distributed/test_finalize_model_grads.py index ee535c29baf..372f8d0d293 100644 --- a/tests/unit_tests/distributed/test_finalize_model_grads.py +++ b/tests/unit_tests/distributed/test_finalize_model_grads.py @@ -11,13 +11,21 @@ from megatron.core.distributed.finalize_model_grads import ( _allreduce_non_tensor_model_parallel_grads, _allreduce_word_embedding_grads, + _update_router_qb_beta, finalize_model_grads, + reset_model_temporary_tensors, +) +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_local_submodules, + get_gpt_layer_with_transformer_engine_spec, ) -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec from megatron.core.models.gpt.gpt_model import GPTModel from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.moe.moe_layer import MoELayer +from megatron.core.transformer.spec_utils import get_submodules from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.training.initialize import _set_random_seed from tests.unit_tests.test_utilities import Utils @@ -117,9 +125,102 @@ def test_finalize_model_grads_requires_custom_group_before_grad_sync(self): assert model.finish_grad_sync_calls == 0 +class TestUpdateRouterQBBeta: + """Exercises the QB bias update in finalize_model_grads against a real MoE router.""" + + def setup_method(self, method): + os.environ.pop('NVTE_FUSED_ATTN', None) + os.environ.pop('NVTE_FLASH_ATTN', None) + os.environ.pop('NVTE_UNFUSED_ATTN', None) + Utils.destroy_model_parallel() + Utils.initialize_model_parallel(1, 1) + _set_random_seed(seed_=123, data_parallel_random_init=False) + + def teardown_method(self, method): + Utils.destroy_model_parallel() + + def _build_moe_layer(self, ema): + num_experts = 8 + config = TransformerConfig( + num_layers=2, + hidden_size=12, + num_attention_heads=4, + num_moe_experts=num_experts, + use_cpu_initialization=True, + moe_router_load_balancing_type="quantile_balancing", + moe_router_score_function="softmax", + moe_router_topk=2, + moe_aux_loss_coeff=0, + moe_router_quantile_balancing_ema=ema, + bf16=True, + params_dtype=torch.bfloat16, + add_bias_linear=False, + ) + submodules = get_submodules( + get_gpt_layer_local_submodules(num_experts=num_experts, moe_grouped_gemm=False).mlp + ) + return config, MoELayer(config, submodules).cuda() + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + @pytest.mark.parametrize("ema", [0.0, 0.9]) + def test_update_router_qb_beta(self, ema): + config, moe_layer = self._build_moe_layer(ema) + router = moe_layer.router + router.train() + # Non-zero prior bias so the EMA term is actually exercised. + router.qb_beta.copy_(torch.randn_like(router.qb_beta)) + + # The real router forward populates qb_beta_accum / qb_beta_count. + hidden = torch.randn((32, 2, config.hidden_size)).cuda().bfloat16() + router(hidden) + router(hidden) + assert router.qb_beta_count.item() == 2 + assert router.qb_beta_accum.abs().sum().item() > 0 + + # Expected from the real accumulators: DP-avg(accum/count), EMA-blend, re-center. + local_avg = router.qb_beta_accum / router.qb_beta_count.clamp(min=1).to(torch.float32) + torch.distributed.all_reduce( + local_avg, op=torch.distributed.ReduceOp.AVG, group=dist.group.WORLD + ) + blended = ema * router.qb_beta + (1.0 - ema) * local_avg + expected = blended - blended.mean(dim=-1, keepdim=True) + + _update_router_qb_beta([moe_layer], config, dp_cp_group=dist.group.WORLD) + + torch.testing.assert_close(router.qb_beta, expected) + torch.testing.assert_close( + router.qb_beta.mean(), torch.zeros((), device=router.qb_beta.device) + ) + + # reset_model_temporary_tensors clears the accumulators for the next global batch. + reset_model_temporary_tensors(config, [moe_layer]) + torch.testing.assert_close(router.qb_beta_accum, torch.zeros_like(router.qb_beta_accum)) + assert router.qb_beta_count.item() == 0 + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + def test_update_router_qb_beta_skips_eval(self): + config, moe_layer = self._build_moe_layer(ema=0.0) + router = moe_layer.router + # Non-zero prior + non-uniform accumulator, so a broken eval guard would visibly + # change qb_beta (a uniform accumulator re-centers to zero and hides the bug). + router.qb_beta.copy_(torch.ones_like(router.qb_beta)) + router.qb_beta_accum.copy_( + torch.arange(router.qb_beta.numel(), dtype=torch.float32, device=router.qb_beta.device) + ) + router.qb_beta_count.fill_(1) + before = router.qb_beta.clone() + router.eval() + + _update_router_qb_beta([moe_layer], config, dp_cp_group=dist.group.WORLD) + + # Eval-mode modules are skipped, so qb_beta is unchanged. + torch.testing.assert_close(router.qb_beta, before) + + class TestAllReduceLNGrads: def init_model(self, share_embeddings_and_output_weights: bool = False): + qk_layernorm = True self.transformer_config = TransformerConfig( num_layers=2, hidden_size=12, @@ -127,13 +228,15 @@ def init_model(self, share_embeddings_and_output_weights: bool = False): use_cpu_initialization=True, tensor_model_parallel_size=self.tp_size, pipeline_model_parallel_size=self.pp_size, - qk_layernorm=True, + qk_layernorm=qk_layernorm, pipeline_dtype=torch.float32, ) self.model = GPTModel( config=self.transformer_config, - transformer_layer_spec=get_gpt_layer_with_transformer_engine_spec(qk_layernorm=True), + transformer_layer_spec=get_gpt_layer_with_transformer_engine_spec( + qk_layernorm=qk_layernorm + ), vocab_size=100, max_sequence_length=4, share_embeddings_and_output_weights=share_embeddings_and_output_weights, diff --git a/tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.py b/tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.py index 6461dbee751..fe56317f39a 100644 --- a/tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.py +++ b/tests/unit_tests/distributed/test_grad_sync_with_expert_parallel.py @@ -8,9 +8,15 @@ from megatron.core import parallel_state from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig +from megatron.core.extensions.transformer_engine import ( + TEColumnParallelGroupedLinear, + TEColumnParallelLinear, +) from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_layer_with_transformer_engine_submodules, ) +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.tensor_parallel.layers import ColumnParallelLinear from megatron.core.transformer import TransformerConfig from megatron.core.transformer.moe.moe_layer import MoELayer, MoESubmodules from megatron.core.transformer.spec_utils import get_submodules @@ -104,6 +110,121 @@ def get_moe_model_and_buffers( ) +def _build_expert_linear(implementation: str, config: TransformerConfig) -> torch.nn.Module: + common_kwargs = { + "input_size": config.hidden_size, + "output_size": config.ffn_hidden_size, + "config": config, + "init_method": config.init_method, + "bias": False, + "skip_bias_add": False, + "is_expert": True, + } + if implementation == "native": + return ColumnParallelLinear( + **common_kwargs, + gather_output=False, + tp_group=parallel_state.get_expert_tensor_parallel_group(), + ) + if implementation == "transformer_engine": + return TEColumnParallelLinear( + **common_kwargs, + gather_output=False, + tp_group=parallel_state.get_expert_tensor_parallel_group(), + ) + if implementation == "transformer_engine_grouped": + return TEColumnParallelGroupedLinear( + num_gemms=config.num_moe_experts, + **common_kwargs, + pg_collection=ProcessGroupCollection.use_mpu_process_groups(), + ) + raise AssertionError(f"Unsupported implementation: {implementation}") + + +@pytest.mark.parametrize( + ("tensor_model_parallel_size", "expert_tensor_parallel_size"), [(2, 1), (1, 2)] +) +@pytest.mark.parametrize( + "implementation", ["native", "transformer_engine", "transformer_engine_grouped"] +) +def test_expert_grad_sync_uses_expert_data_parallel_group( + implementation: str, tensor_model_parallel_size: int, expert_tensor_parallel_size: int +): + """Expert gradients must not be reduced over ordinary DP when ETP differs from TP.""" + if Utils.world_size < 4 or Utils.world_size % 4 != 0: + pytest.skip("Test requires a world size divisible by four") + if Utils.world_size > 16: + pytest.skip("Rank-encoded gradients are intended for small unit-test world sizes") + + Utils.initialize_model_parallel( + tensor_model_parallel_size=tensor_model_parallel_size, + expert_model_parallel_size=1, + expert_tensor_parallel_size=expert_tensor_parallel_size, + ) + try: + # Per-token loss leaves DDP's pre-collective gradient scaling at one. + config = TransformerConfig( + num_layers=1, + hidden_size=8, + num_attention_heads=4, + ffn_hidden_size=16, + num_moe_experts=2, + moe_ffn_hidden_size=16, + moe_router_topk=2, + tensor_model_parallel_size=tensor_model_parallel_size, + expert_model_parallel_size=1, + expert_tensor_parallel_size=expert_tensor_parallel_size, + calculate_per_token_loss=True, + gradient_accumulation_fusion=False, + perform_initialization=False, + bf16=True, + params_dtype=torch.bfloat16, + add_bias_linear=False, + ) + module = _build_expert_linear(implementation, config).cuda() + model = DistributedDataParallel( + config, + ddp_config=DistributedDataParallelConfig( + grad_reduce_in_fp32=True, + overlap_grad_reduce=False, + use_distributed_optimizer=False, + average_in_collective=False, + ), + module=module, + ) + + expert_dp_group = parallel_state.get_expert_data_parallel_group( + partial_expert_data_parallel=True + ) + ordinary_dp_group = parallel_state.get_data_parallel_group( + with_context_parallel=True, partial_data_parallel=True + ) + expert_dp_ranks = torch.distributed.get_process_group_ranks(expert_dp_group) + ordinary_dp_ranks = torch.distributed.get_process_group_ranks(ordinary_dp_group) + assert expert_dp_ranks != ordinary_dp_ranks + + # Powers of two give every rank set a distinct sum, exposing the wrong collective group. + rank_value = float(2 ** torch.distributed.get_rank()) + expected_value = float(sum(2**rank for rank in expert_dp_ranks)) + ordinary_dp_value = float(sum(2**rank for rank in ordinary_dp_ranks)) + assert expected_value != ordinary_dp_value + + for param in model.parameters(): + param.main_grad.fill_(rank_value) + model.finish_grad_sync() + + for param in model.parameters(): + torch.testing.assert_close( + param.main_grad, torch.full_like(param.main_grad, expected_value), rtol=0, atol=0 + ) + + assert not model.buffers + assert len(model.expert_parallel_buffers) == 1 + assert all(param.allreduce is False for param in model.parameters()) + finally: + Utils.destroy_model_parallel() + + @pytest.mark.parametrize("use_distributed_optimizer", [False, True]) @pytest.mark.parametrize("overlap_grad_reduce", [False, True]) @pytest.mark.parametrize("average_in_collective", [False, True]) diff --git a/tests/unit_tests/distributed/test_param_and_grad_buffer.py b/tests/unit_tests/distributed/test_param_and_grad_buffer.py index 281e156efea..4e9c6a27d31 100644 --- a/tests/unit_tests/distributed/test_param_and_grad_buffer.py +++ b/tests/unit_tests/distributed/test_param_and_grad_buffer.py @@ -79,12 +79,19 @@ def get_model_and_buffers( # Wrap with DistributedDataParallel, and get underlying buffer. # Use dummy TransformerConfig with mostly default values. Avoid divide-by-zero # errors for num_attention_heads and num_layers. - # Pre-compute parameter layouts for the distributed optimizer. + # Pre-compute parameter layouts for the distributed optimizer. Size the layout by the group + # the optimizer shards over, which is the intra-instance group when there are several + # optimizer instances. This is the same group DDP hands to the buffer below. full_param_layout = None if use_distributed_optimizer: all_params = [p for p in model.parameters() if p.requires_grad] full_param_layout = DistributedOptimizer.compute_full_param_layout( - all_params, bucket_size, parallel_state.get_data_parallel_world_size(), ddp_config + all_params, + bucket_size, + parallel_state.get_data_parallel_world_size( + with_context_parallel=True, partial_data_parallel=True + ), + ddp_config, ) model = DistributedDataParallel( TransformerConfig(num_attention_heads=1, num_layers=1), @@ -937,6 +944,55 @@ def test_nvfp4_varied_param_sizes(self): assert buffer.param_index_map[params[0]] == (small_unpacked_start, small_unpacked_end, 0) +@pytest.mark.parametrize("num_distributed_optimizer_instances", [1, 2]) +def test_optimizer_shards_cover_every_param(num_distributed_optimizer_instances: int): + """Every parameter must be owned by exactly one rank of the buffer's data-parallel group. + + ``DistributedOptimizer._build_model_gbuf_range`` splits each bucket into N shards and gives + rank ``r`` the r-th one, while the reduce-scatter/all-gather runs over the buffer's + ``data_parallel_group``. N must therefore equal that group's size. If N is larger (e.g. taken + from a layout sized by the full DP world while the group is intra-optimizer-instance), the + trailing shards belong to no rank: those params are never updated by the optimizer and vanish + from grad-norm, num-zeros and params-norm, which are summed over owned shards only. + """ + Utils.initialize_model_parallel( + num_distributed_optimizer_instances=num_distributed_optimizer_instances + ) + + _, param_and_grad_buffer, _ = get_model_and_buffers( + input_dim=100, + output_dim=100, + num_layers=2, + bias=True, + shared_embedding=False, + bucket_size=None, + use_distributed_optimizer=True, + overlap_grad_reduce=False, + average_in_collective=False, + num_distributed_optimizer_instances=num_distributed_optimizer_instances, + ) + + # Sum the parameter elements this rank owns across every bucket. + owned_numel = 0 + for bucket_index in range(len(param_and_grad_buffer.buckets)): + param_map = DistributedOptimizer._build_model_gbuf_range( + param_and_grad_buffer, bucket_index + )["param_map"] + for param_ranges in param_map.values(): + owned_numel += param_ranges["param"].size + + owned_total = torch.tensor([owned_numel], dtype=torch.long, device='cuda') + torch.distributed.all_reduce(owned_total, group=param_and_grad_buffer.data_parallel_group) + + expected_numel = sum(param.numel() for param in param_and_grad_buffer.params) + assert owned_total.item() == expected_numel, ( + f"Optimizer shards cover {owned_total.item()} of {expected_numel} param elements; " + f"{expected_numel - owned_total.item()} elements are owned by no rank" + ) + + Utils.destroy_model_parallel() + + @pytest.mark.parametrize("use_distributed_optimizer", [False, True]) def test_expert_parallel_params_get_separate_buffers(use_distributed_optimizer: bool): """Verify that expert-parallel params (allreduce=False) land in separate buffers @@ -968,7 +1024,12 @@ def test_expert_parallel_params_get_separate_buffers(use_distributed_optimizer: if use_distributed_optimizer: all_params = [p for p in model.parameters() if p.requires_grad] full_param_layout = DistributedOptimizer.compute_full_param_layout( - all_params, bucket_size, parallel_state.get_data_parallel_world_size(), ddp_config + all_params, + bucket_size, + parallel_state.get_data_parallel_world_size( + with_context_parallel=True, partial_data_parallel=True + ), + ddp_config, ) ddp_model = DistributedDataParallel( diff --git a/tests/unit_tests/distributed/test_torch_fully_sharded_parallel.py b/tests/unit_tests/distributed/test_torch_fully_sharded_parallel.py index 6a50e8d1aa5..e1aa1a60b78 100644 --- a/tests/unit_tests/distributed/test_torch_fully_sharded_parallel.py +++ b/tests/unit_tests/distributed/test_torch_fully_sharded_parallel.py @@ -11,14 +11,36 @@ init_num_microbatches_calculator, unset_num_microbatches_calculator, ) +from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel import ColumnParallelLinear from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import MegatronModule +from megatron.core.transformer.mlp import apply_swiglu_sharded_factory from megatron.core.transformer.module import Float16Module from megatron.core.transformer.transformer_config import TransformerConfig -from megatron.core.utils import init_method_normal, is_torch_min_version +from megatron.core.utils import ( + init_method_normal, + is_torch_min_version, + make_tp_sharded_tensor_for_checkpoint, +) from tests.unit_tests.test_utilities import Utils +try: + from torch.distributed import DeviceMesh + from torch.distributed.tensor import DTensor + from torch.distributed.tensor.placement_types import Shard + + HAVE_DTENSOR = True +except ImportError: + HAVE_DTENSOR = False + +try: + import einops + + HAVE_EINOPS = True +except ImportError: + HAVE_EINOPS = False + class DummyModel(MegatronModule): """Setup a few modules to test the FSDP2 constructor.""" @@ -106,3 +128,117 @@ def _is_fsdp_wrapped_module(instance): assert _is_fsdp_wrapped_module(fsdp_model.module.module.linear) assert _is_fsdp_wrapped_module(fsdp_model.module.module.column_parallel_linear) assert not _is_fsdp_wrapped_module(fsdp_model.module.module.conv) + + +@pytest.mark.skipif(not is_torch_min_version("2.4.0"), reason="FSDP2 requires PyTorch >= 2.4") +@pytest.mark.skipif(not HAVE_EINOPS, reason="einops is not available") +@pytest.mark.skipif(not HAVE_DTENSOR, reason="DTensor is not available") +def test_fsdp2_swiglu_sharded_tensor_factory(): + """ + Test construction of a TP2 DP{N} ShardedTensor for SwiGLU. + """ + # Initialize distributed with TP2. + Utils.initialize_model_parallel(tensor_model_parallel_size=2, pipeline_model_parallel_size=1) + pg_collection = ProcessGroupCollection.use_mpu_process_groups() + tp_group = pg_collection.tp + tp_size = tp_group.size() + tp_rank = tp_group.rank() + dp_cp_group = pg_collection.dp_cp + dp_cp_size = dp_cp_group.size() + dp_cp_rank = dp_cp_group.rank() + + # Create FSDP2 DTensor with DP only. (Implicitly TP-sharded before FSDP2 init.) + device_mesh = DeviceMesh.from_group(dp_cp_group, device_type="cuda", mesh_dim_names=["dp_cp"]) + toy_dtensor = DTensor.from_local( + torch.randn(2, 8), device_mesh=device_mesh, placements=(Shard(dim=0),) + ) + toy_dtensor.is_torch_fsdp2_param = True + + # Initialize TP-DP ShardedTensor. + tp_dp_sh_ten = make_tp_sharded_tensor_for_checkpoint( + toy_dtensor, + "test_fsdp2_tp_swiglu_weight", + # SwiGLU FC1 TP & FSDP2 Sharding Dim + tp_axis=0, + replica_id=None, + prepend_offsets=(), + tp_group=tp_group, + dp_cp_group=dp_cp_group, + ) + """ + Before TP2-DP4 Swizzle (TP Rank 1, DP Rank 1): + (Pdb) tp_dp_sh_ten + ShardedTensor( + local_shape=(2, 8), + global_shape=(16, 8), + # Canonical Data Offset = 2 * (tp_rank * dp_cp_size + dp_cp_rank) + # = 2 * (5) = 10 + global_offset=(10, 0), + axis_fragmentations=(8, 1), + replica_id=(0, 0, 0), + prepend_axis_num=0 + ) + """ + + # Test SwiGLU factory for FSDP2-TP. + swiglu_sh_ten_factory = apply_swiglu_sharded_factory( + tp_dp_sh_ten, + sharded_offsets=(), # Vanilla MLP. + # Fused W/V. + singleton_local_shards=False, + tp_group=tp_group, + dp_group=dp_cp_group, + ) + sh_ten_shards = swiglu_sh_ten_factory.build() + + """ + After TP2-DP4 Swizzle (TP Rank 1, DP Rank 1): + (Pdb) sh_ten_shards[0] + ShardedTensor( + local_shape=(2, 8), + global_shape=(16, 8), + global_offset=(6, 0), # W/V TP-swizzled, DP-sharded! + axis_fragmentations=(8, 1), + replica_id=(0, 0, 0), + prepend_axis_num=0 + ) + + This is a mapping from the checkpoint [W;V] rank offsets: + + W_tp0_dp0 W_tp0_dp1 W_tp1_dp0 W_tp1_dp1 V_tp0_dp2 V_tp0_dp3 V_tp1_dp2 V_tp1_dp3 + 0 1 2 3 4 5 6 7 + | + Data Offset = 3 * 2 = 6 is just the rank offset x local shape. + + to the model [ {W_tpx; V_tpx}_dpy ] rank offsets: + + W_tp0_dp0 W_tp0_dp1 V_tp0_dp2 V_tp0_dp3 W_tp1_dp0 W_tp1_dp1 V_tp1_dp2 V_tp1_dp3 + """ + + # Validate FSDP2 TP-DP sharding and swizzle. + assert getattr(tp_dp_sh_ten, "is_torch_fsdp2_param", False) + assert len(sh_ten_shards) == 1 + shard = sh_ten_shards[0] + toy_tensor_shape = toy_dtensor.to_local().shape + assert shard.axis_fragmentations[0] == tp_size * dp_cp_size + assert shard.global_shape[0] == tp_size * dp_cp_size * toy_tensor_shape[0] + # Expected global data offsets considering the parallelism ranks and tensor shape. + expected_global_rank_offsets = { + # (TP Rank, DP Rank) -> Global Data Rank Location / Offset + (0, 0): 0, + (0, 1): 1, + (0, 2): 4, + (0, 3): 5, + (1, 0): 2, + (1, 1): 3, + (1, 2): 6, + (1, 3): 7, + } + assert ( + shard.global_offset[0] + == expected_global_rank_offsets[(tp_rank, dp_cp_rank)] * toy_tensor_shape[0] + ) + + # Destroy distributed. + Utils.destroy_model_parallel() + unset_num_microbatches_calculator() diff --git a/tests/unit_tests/generalized_tensor_parallel/__init__.py b/tests/unit_tests/generalized_tensor_parallel/__init__.py new file mode 100644 index 00000000000..b5dff7b5663 --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/__init__.py @@ -0,0 +1 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. diff --git a/tests/unit_tests/generalized_tensor_parallel/gtp_test_utils.py b/tests/unit_tests/generalized_tensor_parallel/gtp_test_utils.py new file mode 100644 index 00000000000..259cc6ed0d5 --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/gtp_test_utils.py @@ -0,0 +1,158 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Shared fixtures and helpers for all GTP unit tests. +""" + +import pytest +import torch +import transformer_engine.pytorch as te +from transformer_engine.pytorch import is_mxfp8_available, is_nvfp4_available +from transformer_engine.pytorch.quantization import FP8GlobalStateManager + +from megatron.core.tensor_parallel.generalized_tensor_parallelism import GTPShardedParam +from tests.unit_tests.test_utilities import Utils + +# --------------------------------------------------------------------------- +# Fixtures (import into each test module so pytest discovers them) +# --------------------------------------------------------------------------- + + +@pytest.fixture(scope="module", autouse=True) +def _torchrun_dist_init(): + """Initialize the torchrun-managed dist group once per module.""" + Utils.initialize_model_parallel() + yield + Utils.destroy_model_parallel() + + +@pytest.fixture(autouse=True) +def reset_fp8_state(): + yield + FP8GlobalStateManager.reset() + + +@pytest.fixture(autouse=True) +def reset_gtp_globals(): + """Reset GTP mutable class-level state between tests.""" + yield + GTPShardedParam._chain_state = {} + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _run_distributed(fn, required_world_size: int, *args) -> None: + """Run ``fn(rank, world_size, port, *args)`` on every torchrun rank. + + ``port`` is unused (dist already initialized by torchrun) but kept so + worker signatures don't need editing. + """ + actual_world_size = torch.distributed.get_world_size() + if actual_world_size != required_world_size: + pytest.skip( + f"Requires world_size={required_world_size}, " + f"got {actual_world_size} (launch with torchrun --nproc-per-node={required_world_size})" + ) + fn(torch.distributed.get_rank(), actual_world_size, None, *args) + + +def _requires_multi_gpu(n: int = 4): + if torch.cuda.device_count() < n: + pytest.skip(f"Requires at least {n} CUDA devices") + + +def _requires_mxfp8(): + available, reason = is_mxfp8_available(return_reason=True) + if not available: + pytest.skip(f"MXFP8 not available: {reason}") + + +def _requires_nvfp4(): + if not is_nvfp4_available(): + pytest.skip("NVFP4 not available (requires compute capability >= 10.0)") + + +def _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype=torch.bfloat16, **kwargs): + """Construct a bias-free GTP-sharded te.Linear on CUDA. + + Mirrors the production integration (extensions/transformer_engine.py): TE has no GTP + construction hooks, so the module is built stock, ``module.gtp_remat_size`` (the forward-gather + gate) is stamped post-init, and the BF16 weight is sliced by ``wrap_module_params_gtp``. + """ + from megatron.core.tensor_parallel.gtp_api import wrap_module_params_gtp + + layer = te.Linear( + in_features=in_f, + out_features=out_f, + bias=False, + params_dtype=dtype, + device="cuda", + **kwargs, + ) + layer.gtp_remat_size = gtp_remat_group.size() + wrap_module_params_gtp(layer, layer.weight_names, gtp_remat_group) + return layer + + +def _make_gtp_remat_grouped_linear( + num_gemms, in_f, out_f, gtp_remat_group, dtype=torch.bfloat16, **kwargs +): + """Construct a bias-free GTP-sharded te.GroupedLinear on CUDA (post-init slice, see + _make_gtp_linear).""" + from megatron.core.tensor_parallel.gtp_api import wrap_module_params_gtp + + layer = te.GroupedLinear( + num_gemms=num_gemms, + in_features=in_f, + out_features=out_f, + bias=False, + params_dtype=dtype, + device="cuda", + **kwargs, + ) + layer.gtp_remat_size = gtp_remat_group.size() + # GroupedLinear exposes per-expert weight0..weight{num_gemms-1} (it no longer declares + # weight_names); build the names here to match attach_gtp_to_presharded_module. + weight_names = [f"weight{idx}" for idx in range(num_gemms)] + wrap_module_params_gtp(layer, weight_names, gtp_remat_group, is_grouped=True) + return layer + + +def _restore_gtp_shards_and_init_main_grad(module, saved_weights, gtp_rank, dtype=torch.bfloat16): + """Load saved full weights into a GTP_remat_size>1 module and prep it for backward. + + GTPShardedParams receive their ``gtp_rank`` axis-0 shard; replicated params get the full + tensor. Then pre-allocate ``main_grad`` on every GTPShardedParam (required before the first + backward). Used by the dense two-phase baseline-vs-GTP tests. + """ + for name, p in module.named_parameters(): + full = saved_weights[name] + if isinstance(p, GTPShardedParam): + shard_size = p.shape[0] + p.data.copy_(full[gtp_rank * shard_size : (gtp_rank + 1) * shard_size]) + else: + p.data.copy_(full) + for p in module.parameters(): + if isinstance(p, GTPShardedParam): + p.main_grad = torch.zeros(p.shape, dtype=dtype, device='cuda') + + +def _assert_loss_trajectories_match(baseline_losses, test_losses, steps, label="gtp_remat"): + """On rank 0: print and assert two per-step loss trajectories match. + + GTP (ZeRO-3-like) reduces grads with a reduce-scatter-sum while the no-GTP baseline + all-reduces; in BF16 these differ only in reduction order, so the trajectories track to + BF16 precision (observed max |diff| ~1e-2 over 10 steps) rather than bitwise. The tolerance + is set for that BF16 noise floor and still trips on any real GTP sharding bug, which diverges + by O(loss). (The sibling grad-correctness test likewise checks ~1e-2-scale grad error.) + """ + assert ( + len(baseline_losses) == len(test_losses) == steps + ), f"loss counts: baseline={len(baseline_losses)} {label}={len(test_losses)} want {steps}" + for step, (lb, lt) in enumerate(zip(baseline_losses, test_losses)): + print(f"Step {step:2d}: baseline={lb:.6f} {label}={lt:.6f}", flush=True) + torch.testing.assert_close( + torch.tensor(test_losses), torch.tensor(baseline_losses), atol=5e-2, rtol=5e-2 + ) diff --git a/tests/unit_tests/generalized_tensor_parallel/test_attention_gtp.py b/tests/unit_tests/generalized_tensor_parallel/test_attention_gtp.py new file mode 100644 index 00000000000..bfa27e9ae46 --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_attention_gtp.py @@ -0,0 +1,232 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Integration tests for GTP + Attention (TransformerLayer) correctness. + +Test groups +----------- +TestAttentionGTPCorrectness - GTP TransformerLayer loss trajectory matches baseline (no-GTP) + over 10 training steps using MXFP8 and Nemotron3-Super proxy + hyperparameters. +""" + +import pytest +import torch +import torch.distributed as dist + +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + +if not HAVE_GTP: + pytest.skip("GTP requires TransformerEngine >= 2.19", allow_module_level=True) + +from transformer_engine.pytorch import fp8_autocast + +from megatron.core.tensor_parallel.generalized_tensor_parallelism import GTPShardedParam +from tests.unit_tests.generalized_tensor_parallel.gtp_test_utils import ( + _assert_loss_trajectories_match, + _restore_gtp_shards_and_init_main_grad, + _run_distributed, + _torchrun_dist_init, + reset_fp8_state, + reset_gtp_globals, +) + +# --------------------------------------------------------------------------- +# Attention GTP_remat correctness: per-step loss trajectory baseline vs GTP_remat=4 +# --------------------------------------------------------------------------- + + +def _worker_attention_gtp_correctness(rank, world_size, port): + """Verify GTP TransformerLayer produces the same per-step loss as a no-GTP baseline. + + Phase 1 — GTP_remat_size=1, DP=4: + All 4 ranks hold the full model and process identical inputs. Gradients + are identical across ranks (no all-reduce needed). Weight update: + param.data -= lr * param.grad + + Phase 2 — GTP_remat_size=4, DP=1: + All linear weights (QKV proj, output proj, MLP fc1/fc2) sharded across + 4 ranks. After backward, wgrad reduce-scatter sums each shard's wgrad: + main_grad[rank_i] = gtp_remat_size * dW[shard_i] + The optimizer divides by gtp_remat_size to recover the per-element gradient: + param.data -= (lr / gtp_remat_size) * param.main_grad + + Both phases use identical initial weights (synced from rank 0 in Phase 1, + restored as shards in Phase 2) and identical step-by-step inputs. + + Nemotron3-Super proxy hyperparameters: + hidden=4096, num_heads=32 (head_dim=128), ffn_hidden_size=16384 (=4xhidden) + MXFP8 alignment with GTP_remat_size=4: + QKV shard: 3x4096/4=3072, 3072%32=0 ✓; proj shard: 4096/4=1024, 1024%32=0 ✓ + fc1 shard: 16384/4=4096, 4096%32=0 ✓; fc2 shard: 4096/4=1024, 1024%32=0 ✓ + """ + from transformer_engine.pytorch.quantization import FP8GlobalStateManager + + from megatron.core import parallel_state as ps + from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec + from megatron.core.process_groups_config import ProcessGroupCollection + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + from megatron.core.transformer.transformer_config import TransformerConfig + + HIDDEN = 4096 + NUM_HEADS = 32 # head_dim = HIDDEN / NUM_HEADS = 128 + FFN_HIDDEN = 16384 # = 4 x HIDDEN (default GPT FFN ratio) + NUM_LAYERS = 2 + SEQ = 32 + BATCH = 1 + LR = 0.01 + STEPS = 10 + dtype = torch.bfloat16 + + def make_config(): + return TransformerConfig( + num_attention_heads=NUM_HEADS, + num_layers=NUM_LAYERS, + hidden_size=HIDDEN, + ffn_hidden_size=FFN_HIDDEN, + add_bias_linear=False, + params_dtype=dtype, + hidden_dropout=0.0, + attention_dropout=0.0, + bias_dropout_fusion=False, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + ) + + def make_transformer_stack(config, pg_collection): + spec = get_gpt_layer_with_transformer_engine_spec() + return torch.nn.ModuleList( + [ + spec.module( + config, spec.submodules, layer_number=i + 1, pg_collection=pg_collection + ) + for i in range(NUM_LAYERS) + ] + ) + + def run_step(layers, x): + with fp8_autocast(enabled=False): + for layer in layers: + x, _ = layer(x, attention_mask=None) + return x.mean() + + # ------------------------------------------------------------------------- + # Phase 1: Baseline — GTP_remat=1 (DP=4) + # ------------------------------------------------------------------------- + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=1 + ) + model_parallel_cuda_manual_seed(42) + + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=['tp', 'cp', 'gtp_remat'] + ) + config = make_config() + layers = make_transformer_stack(config, pg_collection) + for layer in layers: + layer.cuda() + + # Verify baseline has no GTP_remat sharding (gtp_remat_size=1 should leave plain parameters). + assert not any( + isinstance(p, GTPShardedParam) for p in layers.parameters() + ), "Baseline GTP_remat_size=1 stack should have no GTPShardedParam" + + # Synchronize weights from rank 0 across all DP ranks. + for p in layers.parameters(): + dist.broadcast(p.data, src=0) + + # Save initial weights; will be used to initialize the GTP_remat model identically. + saved_weights = {n: p.data.clone() for n, p in layers.named_parameters()} + + baseline_losses = [] + for step in range(STEPS): + torch.manual_seed(step) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + dist.broadcast(x, src=0) + + loss = run_step(layers, x) + if rank == 0: + baseline_losses.append(loss.item()) + + loss.backward() + with torch.no_grad(): + for p in layers.parameters(): + if p.grad is not None: + p.data.sub_(LR * p.grad) + p.grad.zero_() + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + FP8GlobalStateManager.reset() + + # ------------------------------------------------------------------------- + # Phase 2: GTP_remat=4 (DP=1) + # ------------------------------------------------------------------------- + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=4 + ) + model_parallel_cuda_manual_seed(42) + + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=['tp', 'cp', 'gtp_remat'] + ) + config = make_config() + layers_gtp = make_transformer_stack(config, pg_collection) + for layer in layers_gtp: + layer.cuda() + + gtp_remat_group = ps.get_gtp_weight_remat_group() + gtp_remat_size = gtp_remat_group.size() + gtp_rank = gtp_remat_group.rank() + + # Verify GTP_remat is truly active: linear weights must be GTPShardedParam instances. + gtp_params = [p for p in layers_gtp.parameters() if isinstance(p, GTPShardedParam)] + assert ( + len(gtp_params) > 0 + ), "GTP is not active: no GTPShardedParam found in GTP_remat_size=4 transformer stack" + + # Restore initial weights into shards and pre-allocate main_grad for the backward. + _restore_gtp_shards_and_init_main_grad(layers_gtp, saved_weights, gtp_rank, dtype) + + gtp_losses = [] + for step in range(STEPS): + for p in layers_gtp.parameters(): + if isinstance(p, GTPShardedParam): + p.main_grad.zero_() + + torch.manual_seed(step) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + dist.broadcast(x, src=0) + + loss = run_step(layers_gtp, x) + if rank == 0: + gtp_losses.append(loss.item()) + + loss.backward() + + # After RS, main_grad = gtp_remat_size * dW_shard. Divide by gtp_remat_size for baseline. + with torch.no_grad(): + for p in layers_gtp.parameters(): + if isinstance(p, GTPShardedParam): + p.data.sub_((LR / gtp_remat_size) * p.main_grad) + elif p.grad is not None: + p.data.sub_(LR * p.grad) + p.grad.zero_() + + ps.destroy_model_parallel() + ps.initialize_model_parallel() + GTPShardedParam._chain_state = {} + + # ------------------------------------------------------------------------- + # Compare per-step loss trajectories on rank 0 + # ------------------------------------------------------------------------- + if rank == 0: + _assert_loss_trajectories_match(baseline_losses, gtp_losses, STEPS) + + +class TestAttentionGTPCorrectness: + def test_attention_gtp_loss_trajectory_matches_baseline(self): + """GTP TransformerLayer per-step losses must match no-GTP baseline (atol=1e-5, rtol=1e-5; MXFP8, Nemotron3-Super proxy).""" + if torch.cuda.device_count() < 4: + pytest.skip("Requires at least 4 CUDA devices") + _run_distributed(_worker_attention_gtp_correctness, 4) diff --git a/tests/unit_tests/generalized_tensor_parallel/test_gtp_basics.py b/tests/unit_tests/generalized_tensor_parallel/test_gtp_basics.py new file mode 100644 index 00000000000..f488d64ae6a --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_gtp_basics.py @@ -0,0 +1,1208 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Unit tests for Generalized Tensor Parallelism (GTP). + +Scope: sharding math, module wiring, and behavior/regression guards. End-to-end fwd/bwd/loss/grad, +fp8, and checkpoint correctness live in the integration tests (test_gtp_loss_correctness, +test_gtp_grad_correctness, test_attention_gtp, test_mamba_gtp, test_moe_egtp, +test_gtp_fp8_param_gather, test_gtp_dcp), so low-level plumbing smoke tests are not duplicated here. + +Test groups +----------- +- TestGTPSharding - wrap_module_params_gtp: shard content + padding +- TestWrapModuleParams - wrap_module_params_gtp: param replacement + weight_list +- TestLinearGTP / TestLayerNormLinearGTP / TestGroupedLinearGTP - single-layer fwd/bwd +- TestGTPPrefetchChain - linked-list next_w/prev_w wiring +- TestGTPWgradRS - wgrad reduce-scatter shape + multi-layer deferred path +- TestGTPMicrobatches - output consistency across microbatches +- TestMXFP8LinearGTP - Linear + MXFP8 recipe: quantized shard setup, fwd/bwd, padding +- TestGTPGroupSizeOne - wrap_module_params_gtp no-op when gtp_remat_group.size()==1 +- TestGTPPrefetchDisabled - weight_prefetch=False single-pass forward +- TestFuseWgradAccumulation - fuse_wgrad_accumulation=True: wgrad -> main_grad +- TestGTPGradAccumHook - main_grad updated after reduce-scatter backward +- TestWaitAsyncCommsFallback - inline-accumulation fallback when _wgrad_rs_handle is None +- TestGTPDDPBucketAlignment - GTP/regular DDP bucket ends padded for dist-opt alignment +- TestGTPDDPGradReadyWiring - GTP params drive DDP grad-ready via the manual hook, not autograd + +Multi-GPU tests skip when ``torch.distributed.get_world_size()`` != the required world size (4). +""" + +import pytest +import torch +import torch.distributed as dist +import torch.nn as nn + +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + +if not HAVE_GTP: + pytest.skip("GTP requires TransformerEngine >= 2.19", allow_module_level=True) + +import transformer_engine.pytorch as te +from transformer_engine.pytorch import fp8_autocast +from transformer_engine.pytorch.quantized_tensor import QuantizedTensor + +import megatron.core.tensor_parallel.generalized_tensor_parallelism as gtp_module +from megatron.core import parallel_state +from megatron.core.process_groups_config import ProcessGroupCollection +from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( + GTPShardedParam, + wrap_module_params_gtp, +) +from tests.unit_tests.generalized_tensor_parallel.gtp_test_utils import ( + _make_gtp_linear, + _make_gtp_remat_grouped_linear, + _requires_multi_gpu, + _requires_mxfp8, + _run_distributed, + _torchrun_dist_init, + reset_fp8_state, + reset_gtp_globals, +) + + +class _FakeGroup: + """Minimal mock for a dist process group — used in single-process unit tests.""" + + def __init__(self, size=1, rank=0): + self._size = size + self._rank = rank + + def size(self): + return self._size + + def rank(self): + return self._rank + + +def _worker_sharding_aligned(rank, world_size, port): + K, M = world_size * 32, 16 # K divisible by 16*world_size → no padding + full_weight = torch.arange(K * M, dtype=torch.float32).reshape(K, M).cuda() + dist.broadcast(full_weight, src=0) + + gtp_remat_group = dist.new_group(list(range(world_size))) + mod = nn.Module() + mod.weight = nn.Parameter(full_weight.clone(), requires_grad=False) + wrap_module_params_gtp(mod, ["weight"], gtp_remat_group) + shard = mod.weight + + rows_per_rank = K // world_size + assert shard.shape == (rows_per_rank, M), f"rank {rank}: unexpected shape {shard.shape}" + assert shard.pad_length == 0 + expected = full_weight[rank * rows_per_rank : (rank + 1) * rows_per_rank] + assert torch.allclose(shard.data, expected), f"rank {rank}: shard content mismatch" + + +def _worker_sharding_padding(rank, world_size, port): + alignment = 16 * world_size + K = alignment - 1 # deliberately unaligned + M = 16 + full_weight = torch.ones(K, M, dtype=torch.float32).cuda() + dist.broadcast(full_weight, src=0) + + gtp_remat_group = dist.new_group(list(range(world_size))) + mod = nn.Module() + mod.weight = nn.Parameter(full_weight.clone(), requires_grad=False) + wrap_module_params_gtp(mod, ["weight"], gtp_remat_group) + shard = mod.weight + + padded_K = alignment + rows_per_rank = padded_K // world_size + + if rank == world_size - 1: + assert shard.pad_length > 0 + # The shard tensor holds only the real rows; get_padded_shard() appends zero rows. + padded = shard.get_padded_shard() + assert ( + padded.shape[0] == rows_per_rank + ), f"rank {rank}: expected padded shard {rows_per_rank} rows, got {padded.shape[0]}" + n_real = K - rank * rows_per_rank + assert torch.all(padded[n_real:] == 0), "Padding rows must be zero" + else: + # pad_length is set globally on every rank's shard (slicer attaches the + # global padding amount), so we don't assert anything about it here — + # only the last rank's shard contains the actual padding rows. + assert ( + shard.shape[0] == rows_per_rank + ), f"rank {rank}: expected {rows_per_rank} rows, got {shard.shape[0]}" + + +class TestGTPSharding: + def test_aligned_shard_content(self): + _requires_multi_gpu(4) + _run_distributed(_worker_sharding_aligned, 4) + + def test_unaligned_shard_padding(self): + _requires_multi_gpu(4) + _run_distributed(_worker_sharding_padding, 4) + + +# --------------------------------------------------------------------------- +# wrap_module_params_gtp: param replacement and GroupedLinear weight_list +# --------------------------------------------------------------------------- + + +def _worker_linear_param_replaced(rank, world_size, port): + in_f, out_f = 64, 128 + gtp_remat_group = dist.new_group(list(range(world_size))) + layer = _make_gtp_linear(in_f, out_f, gtp_remat_group) + w = layer.weight + assert isinstance(w, GTPShardedParam), "weight must be GTPShardedParam" + assert w.shape == (out_f // world_size, in_f), f"unexpected shard shape {w.shape}" + assert w.group is gtp_remat_group + + +def _worker_grouped_weight_list(rank, world_size, port): + num_gemms, in_f, out_f = 3, 32, 64 + gtp_remat_group = dist.new_group(list(range(world_size))) + layer = _make_gtp_remat_grouped_linear(num_gemms, in_f, out_f, gtp_remat_group) + w0 = layer.weight0 + assert isinstance(w0, GTPShardedParam) + assert w0.weight_list is not None + assert len(w0.weight_list) == num_gemms + assert [w.expert_idx for w in w0.weight_list] == list(range(num_gemms)) + + +class TestWrapModuleParams: + def test_linear_weight_replaced(self): + _requires_multi_gpu(4) + _run_distributed(_worker_linear_param_replaced, 4) + + def test_grouped_linear_weight_list(self): + _requires_multi_gpu(4) + _run_distributed(_worker_grouped_weight_list, 4) + + +# --------------------------------------------------------------------------- +# Linear forward/backward numerical correctness +# --------------------------------------------------------------------------- + + +def _worker_linear_correctness(rank, world_size, port): + """GTP output == (all-gathered weight) @ input, and dX matches.""" + torch.manual_seed(0) + batch, in_f, out_f = 16, 64, 128 # out_f % (16*world_size)==0 → no padding + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + layer = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + + # Reconstruct full weight from shards (all-gather) + shard = layer.weight.data.clone() + all_shards = [torch.zeros_like(shard) for _ in range(world_size)] + dist.all_gather(all_shards, shard, group=gtp_remat_group) + full_weight = torch.cat(all_shards, dim=0).float()[:out_f] # strip any padding + + # Shared input across ranks + inp = torch.randn(batch, in_f, dtype=dtype, device="cuda") + dist.broadcast(inp, src=0) + + inp_gtp = inp.clone().requires_grad_(True) + inp_ref = inp.clone().requires_grad_(True) + + # GTP_remat forward + out_gtp = layer(inp_gtp, is_first_microbatch=True) + + # Reference forward + out_ref = inp_ref.float() @ full_weight.T + out_ref = out_ref.to(dtype) + + assert out_gtp.shape == out_ref.shape, f"Shape mismatch {out_gtp.shape} vs {out_ref.shape}" + assert torch.allclose( + out_gtp.float(), out_ref.float(), atol=1e-5, rtol=1e-5 + ), f"Output mismatch max_diff={(out_gtp.float()-out_ref.float()).abs().max():.4f}" + + # wgrad RS path always accumulates into main_grad; allocate before backward. + layer.weight.main_grad = torch.zeros(layer.weight.shape, dtype=dtype, device="cuda") + + # Backward: compare input gradient + grad_out = torch.randn_like(out_gtp) + dist.broadcast(grad_out, src=0) + out_gtp.backward(grad_out) + out_ref.backward(grad_out.float()) + + assert inp_gtp.grad is not None + assert torch.allclose( + inp_gtp.grad.float(), inp_ref.grad.float(), atol=1e-5, rtol=1e-5 + ), f"dX mismatch max_diff={(inp_gtp.grad.float()-inp_ref.grad.float()).abs().max():.4f}" + + +class TestLinearGTP: + def test_forward_backward_correctness(self): + _requires_multi_gpu(4) + _run_distributed(_worker_linear_correctness, 4) + + +# --------------------------------------------------------------------------- +# LayerNormLinear forward/backward smoke test +# --------------------------------------------------------------------------- + + +def _worker_layernorm_linear(rank, world_size, port): + torch.manual_seed(0) + seq, batch, in_f, out_f = 4, 2, 64, 128 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + layer = te.LayerNormLinear( + in_features=in_f, out_features=out_f, bias=False, params_dtype=dtype, device="cuda" + ) + # TE construction is GTP-agnostic: gtp_remat_size (forward-gather gate) is stamped and the + # BF16 weight is sliced post-init (Megatron side). + layer.gtp_remat_size = gtp_remat_group.size() + wrap_module_params_gtp(layer, layer.weight_names, gtp_remat_group) + assert isinstance(layer.weight, GTPShardedParam) + + inp = torch.randn(seq, batch, in_f, dtype=dtype, device="cuda", requires_grad=True) + dist.broadcast(inp, src=0) + + out = layer(inp, is_first_microbatch=True) + assert out.shape == (seq, batch, out_f), f"unexpected output shape {out.shape}" + + layer.weight.main_grad = torch.zeros(layer.weight.shape, dtype=dtype, device="cuda") + out.sum().backward() + assert inp.grad is not None and inp.grad.shape == inp.shape + + +class TestLayerNormLinearGTP: + def test_forward_backward(self): + _requires_multi_gpu(4) + _run_distributed(_worker_layernorm_linear, 4) + + +# --------------------------------------------------------------------------- +# GroupedLinear forward/backward smoke test +# --------------------------------------------------------------------------- + + +def _worker_grouped_linear(rank, world_size, port, num_gemms): + torch.manual_seed(0) + in_f, out_f, total_tokens = 32, 64, num_gemms * 4 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + layer = _make_gtp_remat_grouped_linear(num_gemms, in_f, out_f, gtp_remat_group, dtype) + assert isinstance(layer.weight0, GTPShardedParam) + + m_splits = [total_tokens // num_gemms] * num_gemms + m_splits[-1] += total_tokens - sum(m_splits) + + inp = torch.randn(total_tokens, in_f, dtype=dtype, device="cuda", requires_grad=True) + dist.broadcast(inp, src=0) + + out = layer(inp, m_splits=m_splits, is_first_microbatch=True) + assert out.shape == (total_tokens, out_f), f"unexpected output shape {out.shape}" + + for i in range(num_gemms): + w = getattr(layer, f"weight{i}") + w.main_grad = torch.zeros(w.shape, dtype=dtype, device="cuda") + out.sum().backward() + assert inp.grad is not None and inp.grad.shape == inp.shape + + +class TestGroupedLinearGTP: + @pytest.mark.parametrize("num_gemms", [2, 4]) + def test_forward_backward(self, num_gemms): + _requires_multi_gpu(4) + _run_distributed(_worker_grouped_linear, 4, num_gemms) + + +def _worker_ops_grouped_linear(rank, world_size, port, num_gemms): + """GTP on the fusible-op ``te.ops.GroupedLinear`` -- the unfused fallback path (run standalone, + so no op fusion). Exercises materialize (fwd/bwd all-gather) + wgrad reduce-scatter wiring in + transformer_engine/pytorch/ops/basic/grouped_linear.py.""" + torch.manual_seed(0) + in_f, out_f, total_tokens = 32, 64, num_gemms * 4 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + op = te.ops.GroupedLinear(num_gemms, in_f, out_f, bias=False, device="cuda", dtype=dtype) + op.gtp_remat_size = gtp_remat_group.size() + wrap_module_params_gtp( + op, [f"weight{i}" for i in range(num_gemms)], gtp_remat_group, is_grouped=True + ) + assert isinstance(op.weight0, GTPShardedParam) + + m_splits = [total_tokens // num_gemms] * num_gemms + m_splits[-1] += total_tokens - sum(m_splits) + split_sizes = torch.tensor(m_splits, dtype=torch.int64, device="cuda") + + inp = torch.randn(total_tokens, in_f, dtype=dtype, device="cuda", requires_grad=True) + dist.broadcast(inp, src=0) + + out = op(inp, split_sizes) + assert out.shape == (total_tokens, out_f), f"unexpected output shape {out.shape}" + + for i in range(num_gemms): + w = getattr(op, f"weight{i}") + w.main_grad = torch.zeros(w.shape, dtype=dtype, device="cuda") + # DDP initializes this on every param; the backward wgrad-fusion path sets it True and + # returns a throwaway dummy .grad (real grad is reduce-scattered into main_grad). + w.grad_added_to_main_grad = False + out.sum().backward() + assert inp.grad is not None and inp.grad.shape == inp.shape + # The wgrad reduce-scatter wrote each per-expert shard's gradient into main_grad (the gradient + # of record for GTP), and flagged it so DDP won't double-add the dummy .grad. + for i in range(num_gemms): + w = getattr(op, f"weight{i}") + assert w.grad_added_to_main_grad is True + assert w.main_grad.shape == w.shape + assert torch.count_nonzero(w.main_grad) > 0, f"weight{i} main_grad not populated by RS" + + +class TestOpsGroupedLinearGTP: + """GTP on the fusible-op ``te.ops.GroupedLinear`` (the unfused fallback for grouped-MLP).""" + + @pytest.mark.parametrize("num_gemms", [2, 4]) + def test_forward_backward(self, num_gemms): + _requires_multi_gpu(4) + _run_distributed(_worker_ops_grouped_linear, 4, num_gemms) + + +# --------------------------------------------------------------------------- +# Prefetch chain: next_w / prev_w wiring after first forward pass +# --------------------------------------------------------------------------- + + +def _worker_chain_wired(rank, world_size, port): + torch.manual_seed(0) + in_f, out_f = 32, 64 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + l0 = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + l1 = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + + inp = torch.randn(4, in_f, dtype=dtype, device="cuda") + dist.broadcast(inp, src=0) + + # First forward pass builds the linked list + l0(inp, is_first_microbatch=True) + l1(inp, is_first_microbatch=True) + + w0, w1 = l0.weight, l1.weight + assert w0.next_w is w1, "w0.next_w should point to w1" + assert w1.prev_w is w0, "w1.prev_w should point back to w0" + assert w1.next_w is None + assert w0.prev_w is None + + +def _worker_chain_async_prefetch(rank, world_size, port): + """On the second forward pass, w1 should be in DATA_READY before its forward runs.""" + torch.manual_seed(0) + in_f, out_f = 32, 64 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + l0 = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + l1 = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + + inp = torch.randn(4, in_f, dtype=dtype, device="cuda") + dist.broadcast(inp, src=0) + + # First pass builds chain, second pass uses async prefetch + for _ in range(2): + out = l0(inp, is_first_microbatch=True) + l1(inp, is_first_microbatch=True) + assert torch.isfinite(out).all(), "Non-finite output on second pass" + + +class TestGTPPrefetchChain: + def test_chain_wired_after_first_pass(self): + _requires_multi_gpu(4) + _run_distributed(_worker_chain_wired, 4) + + def test_async_prefetch_second_pass(self): + _requires_multi_gpu(4) + _run_distributed(_worker_chain_async_prefetch, 4) + + +class TestGroupedExpertChainClassification: + """Routed grouped experts get their own per-role homogeneous prefetch chains + (one-block-ahead), while sharing a single IB stream. Pure classification logic, + no GPU/distributed needed.""" + + FC1 = "decoder.layers.3.mlp.experts.linear_fc1.weight0" + FC2 = "decoder.layers.3.mlp.experts.linear_fc2.weight0" + SHARED = "decoder.layers.3.mlp.shared_experts.linear_fc1.weight" + MIXER = "decoder.layers.3.mixer.in_proj.weight" + + def teardown_method(self, method): + # Restore the module default so other tests see a clean CG config. + gtp_module.set_cuda_graph_modules(None, cuda_graph_impl="none") + + def test_ungraphed_moe_splits_fc1_fc2_into_own_chains(self): + gtp_module.set_cuda_graph_modules(None, cuda_graph_impl="none") + c1 = gtp_module._classify_param_chain(self.FC1) + c2 = gtp_module._classify_param_chain(self.FC2) + assert c1 == "GTP_remat_grouped_fc1_ungraphed", c1 + assert c2 == "GTP_remat_grouped_fc2_ungraphed", c2 + # Separate linked-list chains so next_w links consecutive MoE layers (one-block-ahead). + assert c1 != c2 + # Removed from the general chain; other layer kinds are unaffected. + assert gtp_module._classify_param_chain(self.SHARED) == "GTP_ungraphed" + assert gtp_module._classify_param_chain(self.MIXER) == "GTP_ungraphed" + + def test_fc1_fc2_share_one_ib_stream(self): + gtp_module.set_cuda_graph_modules(None, cuda_graph_impl="none") + group = object() # same EGTP group for both roles + c1 = gtp_module._classify_param_chain(self.FC1) + c2 = gtp_module._classify_param_chain(self.FC2) + # Distinct chains but ONE shared IB stream (serialize fc1 then fc2). + assert gtp_module._stream_key(c1, group) == gtp_module._stream_key(c2, group) + # Distinct from the general ungraphed chain's stream. + assert gtp_module._stream_key(c1, group) != gtp_module._stream_key("GTP_ungraphed", group) + + def test_graphed_moe_keeps_grouped_in_plain_graphed_chain(self): + # When MoE is captured, grouped weights stay in the plain GRAPHED chain so the + # cross-graph drain wait_async_comms(GRAPHED) still targets them by exact chain id. + gtp_module.set_cuda_graph_modules({"moe"}, cuda_graph_impl="local") + assert gtp_module._classify_param_chain(self.FC1) == "GTP_graphed" + assert gtp_module._classify_param_chain(self.FC2) == "GTP_graphed" + + def test_graphness_helpers(self): + # "ungraphed" is the eager suffix; everything else is captured. + assert not gtp_module._chain_is_graphed("GTP_remat_grouped_fc1_ungraphed") + assert not gtp_module._chain_is_graphed("GTP_ungraphed") + assert gtp_module._chain_is_graphed("GTP_graphed") + + +class TestGroupedDoubleBuffer: + """One-block-ahead grouped chains must double-buffer: consecutive MoE layers get distinct + gather buffers (else prefetching layer N+1 clobbers layer N's in-use weight). Pure cache-key + logic, no GPU/distributed needed.""" + + class _Fake: + _unsharded_shape_padded = (128, 256) + expert_idx = 0 + + def __init__(self, chain_id): + self.chain_id = chain_id + + _double_buffer_parity = gtp_module.GTPShardedParam._double_buffer_parity + _get_cache_key = gtp_module.GTPShardedParam._get_cache_key + + def setup_method(self, method): + gtp_module.reset_gtp_state() + + def teardown_method(self, method): + gtp_module.reset_gtp_state() + + def _key(self, chain_id): + return self._Fake(chain_id)._get_cache_key(torch.bfloat16, fwd=True, reduce_scatter=False) + + def test_consecutive_layers_use_two_alternating_buffers(self): + keys = [self._key("GTP_remat_grouped_fc1_ungraphed") for _ in range(4)] + # Consecutive layers differ (no clobber); alternating layers share; exactly two buffers. + assert keys[0] != keys[1] + assert keys[1] != keys[2] + assert keys[0] == keys[2] + assert keys[1] == keys[3] + assert len(set(keys)) == 2 + + def test_fc1_fc2_never_share_a_buffer(self): + # Both can be in-flight at once on the shared IB stream; role folded into key keeps + # them distinct even when gathered shapes match (as in this fake). + assert self._key("GTP_remat_grouped_fc1_ungraphed") != self._key( + "GTP_remat_grouped_fc2_ungraphed" + ) + + def test_non_grouped_key_unchanged(self): + assert self._key("GTP_ungraphed") == ((128, 256), torch.bfloat16, 0, False) + + def test_parity_cached_and_stable(self): + f = self._Fake("GTP_remat_grouped_fc1_ungraphed") + fwd = f._get_cache_key(torch.bfloat16, fwd=True, reduce_scatter=False) + bwd = f._get_cache_key(torch.bfloat16, fwd=False, reduce_scatter=False) + rs = f._get_cache_key(torch.bfloat16, fwd=False, reduce_scatter=True) + # Same parity for all of this weight's buffers (distinct from neighbours, consistent here). + assert f._buf_parity == 0 + assert fwd[-1] == 0 and bwd[-1] == 0 and rs[-1] == 0 + + +# --------------------------------------------------------------------------- +# Wgrad reduce-scatter: shape and deferred async path +# --------------------------------------------------------------------------- + + +def _worker_wgrad_shape(rank, world_size, port): + """After backward, weight.grad shape must match the local shard shape.""" + torch.manual_seed(0) + in_f, out_f = 32, 64 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + layer = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype, fuse_wgrad_accumulation=False) + inp = torch.randn(8, in_f, dtype=dtype, device="cuda", requires_grad=True) + dist.broadcast(inp, src=0) + + layer.weight.main_grad = torch.zeros(layer.weight.shape, dtype=dtype, device="cuda") + layer(inp, is_first_microbatch=True).sum().backward() + + w = layer.weight + if w.grad is not None: + assert w.grad.shape == w.shape, f"wgrad shape {w.grad.shape} != shard shape {w.shape}" + + +def _worker_multilayer_deferred_rs(rank, world_size, port): + """Two-layer GTP: async RS deferred for layer0 (non-last), sync for layer1 (last in bwd).""" + torch.manual_seed(0) + in_f, out_f = 32, 64 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + l0 = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + l1 = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + + inp = torch.randn(8, in_f, dtype=dtype, device="cuda", requires_grad=True) + dist.broadcast(inp, src=0) + + # wgrad RS path always accumulates into main_grad; allocate before backward. + l0.weight.main_grad = torch.zeros(l0.weight.shape, dtype=dtype, device="cuda") + l1.weight.main_grad = torch.zeros(l1.weight.shape, dtype=dtype, device="cuda") + + out = l0(inp, is_first_microbatch=True) + l1(inp, is_first_microbatch=True) + out.sum().backward() + + # Both weights' main_grad should have been updated + for lyr in [l0, l1]: + w = lyr.weight + assert w.main_grad is not None, f"No main_grad on {lyr.__class__.__name__}.weight" + + +class TestGTPWgradRS: + def test_wgrad_shape_matches_shard(self): + _requires_multi_gpu(4) + _run_distributed(_worker_wgrad_shape, 4) + + def test_multilayer_deferred_rs(self): + _requires_multi_gpu(4) + _run_distributed(_worker_multilayer_deferred_rs, 4) + + +# --------------------------------------------------------------------------- +# Multiple microbatches: output must be consistent when weight unchanged +# --------------------------------------------------------------------------- + + +def _worker_microbatches(rank, world_size, port): + torch.manual_seed(0) + batch, in_f, out_f = 8, 64, 128 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + layer = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + inp = torch.randn(batch, in_f, dtype=dtype, device="cuda") + dist.broadcast(inp, src=0) + + # First microbatch + out1 = layer(inp, is_first_microbatch=True).detach().clone() + + # Second microbatch with same weight (skip_weight_cast=True path) + out2 = layer(inp, is_first_microbatch=False).detach() + + assert torch.allclose( + out1, out2 + ), f"Microbatch outputs differ; max_diff={(out1-out2).abs().max():.6f}" + + +class TestGTPMicrobatches: + def test_consistent_across_microbatches(self): + _requires_multi_gpu(4) + _run_distributed(_worker_microbatches, 4) + + +# --------------------------------------------------------------------------- +# MXFP8 + GTP_remat: Linear forward/backward, quantized shard setup +# --------------------------------------------------------------------------- + + +def _make_native_fp8_gtp_linear(in_f, out_f, gtp_remat_group, dtype, recipe): + """Build a native-FP8 GTP te.Linear the gtp-agnostic way. + + Mirrors megatron/core/extensions/transformer_engine.py: pass the pre-sharded + out_features to a STOCK te.Linear under fp8_model_init (TE inits+quantizes a native + MXFP8 shard with no GTP awareness), then attach the GTP wiring post-init. + """ + from transformer_engine.pytorch import fp8_model_init + + from megatron.core.tensor_parallel.gtp_api import ( + attach_gtp_to_presharded_module, + gtp_remat_shard_dim0, + ) + + shard_out, pad_length = gtp_remat_shard_dim0(out_f, gtp_remat_group) + with fp8_model_init(enabled=True, recipe=recipe): + layer = te.Linear( + in_features=in_f, out_features=shard_out, bias=False, params_dtype=dtype, device="cuda" + ) + layer.gtp_remat_size = gtp_remat_group.size() + attach_gtp_to_presharded_module(layer, gtp_remat_group, pad_length) + return layer + + +def _worker_mxfp8_linear(rank, world_size, port): + """Verify GTP Linear with a native MXFP8 param: all-gather + GEMM + backward. + + mxfp8 always implies --fp8-param-gather: the weight is built as a native FP8 shard at + construction (no BF16 source, no per-forward cast). + """ + from transformer_engine.common.recipe import MXFP8BlockScaling + + from megatron.core.tensor_parallel.generalized_tensor_parallelism import update_gtp_config + + torch.manual_seed(0) + # batch=32: MXFP8 wgrad GEMM (K=batch) requires K divisible by MXFP8_BLOCK_SCALING_SIZE=32 + batch, in_f, out_f = 32, 64, 128 # out_f % (16*world_size)==0 → no padding + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + recipe = MXFP8BlockScaling() + layer = _make_native_fp8_gtp_linear(in_f, out_f, gtp_remat_group, dtype, recipe) + + # The weight IS the native FP8 shard: a QuantizedTensor with the GTP surface attached. + w = layer.weight + assert isinstance(w, QuantizedTensor), f"weight should be QuantizedTensor, got {type(w)}" + assert w.quantized is w, "native-FP8 GTP: self.quantized must be the param itself" + assert getattr(w, "is_gtp_weight_remat", False), "GTP surface missing on native param" + assert w.shape[0] * world_size == out_f, "weight must be dim-0 sharded" + + inp = torch.randn(batch, in_f, dtype=dtype, device="cuda", requires_grad=True) + dist.broadcast(inp, src=0) + + with fp8_autocast(enabled=True, fp8_recipe=recipe): + out = layer(inp, is_first_microbatch=True) + + assert out.shape == (batch, out_f), f"unexpected output shape {out.shape}" + assert torch.isfinite(out).all(), "MXFP8 GTP output has non-finite values" + + # Backward should complete without error + layer.weight.main_grad = torch.zeros(layer.weight.shape, dtype=dtype, device="cuda") + out.sum().backward() + assert inp.grad is not None + assert inp.grad.shape == inp.shape + + # Second microbatch reuses the same native FP8 weight + with fp8_autocast(enabled=True, fp8_recipe=recipe): + out2 = layer(inp.detach(), is_first_microbatch=False) + assert torch.isfinite(out2).all(), "MXFP8 GTP second-microbatch output has non-finite" + + +def _worker_mxfp8_linear_unaligned(rank, world_size, port): + """Verify native-FP8 MXFP8 GTP when out_features needs padding. + + MXFP8 requires tensor dims divisible by 32, so shard_size (= M_padded / world_size) + must be a multiple of 32. With world_size=4 this requires M_padded % 128 == 0. + out_f=120 gives M_padded=128, shard_size=32 (32 % 32 == 0). The last rank's shard + holds 24 real rows zero-padded to 32. After all-gather, _strip_padding removes the + padded rows before the GEMM, so the output has the original out_f columns. + """ + from transformer_engine.common.recipe import MXFP8BlockScaling + + from megatron.core.tensor_parallel.generalized_tensor_parallelism import update_gtp_config + + torch.manual_seed(0) + # out_f=120: M_padded=128, shard_size=32, last rank has 24 rows padded to 32. + out_f = 120 + in_f = 64 + batch = 32 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + recipe = MXFP8BlockScaling() + layer = _make_native_fp8_gtp_linear(in_f, out_f, gtp_remat_group, dtype, recipe) + assert layer.weight.pad_length == 8, f"expected pad 8, got {layer.weight.pad_length}" + + inp = torch.randn(batch, in_f, dtype=dtype, device="cuda", requires_grad=True) + dist.broadcast(inp, src=0) + + with fp8_autocast(enabled=True, fp8_recipe=recipe): + out = layer(inp, is_first_microbatch=True) + + # After _strip_padding removes the padded rows, output has out_f (not padded) cols. + assert out.shape == (batch, out_f), f"unexpected output shape {out.shape}" + assert torch.isfinite(out).all(), "MXFP8 GTP (unaligned) output has non-finite values" + + +class TestMXFP8LinearGTP: + def test_forward_backward(self): + _requires_mxfp8() + _requires_multi_gpu(4) + _run_distributed(_worker_mxfp8_linear, 4) + + def test_forward_unaligned_padding(self): + _requires_mxfp8() + _requires_multi_gpu(4) + _run_distributed(_worker_mxfp8_linear_unaligned, 4) + + +# --------------------------------------------------------------------------- +# wrap_module_params_gtp is a no-op when gtp_remat_group.size() == 1 +# --------------------------------------------------------------------------- + + +class TestGTPGroupSizeOne: + + def test_no_sharding_when_gtp_remat_size_one(self): + """wrap_module_params_gtp must be a no-op for a singleton GTP group.""" + mod = nn.Linear(32, 64, bias=False) + original_weight = mod.weight + wrap_module_params_gtp(mod, ["weight"], _FakeGroup()) + assert ( + mod.weight is original_weight + ), "gtp_remat_group.size()==1 should leave parameters unchanged" + assert not isinstance(mod.weight, GTPShardedParam) + + +class TestGTPRematPgCollectionWithoutParallelState: + """Resolving the GTP shard group with GTP off must return None, not assert. + + TE linear ``__init__`` calls ``use_mpu_process_groups(["gtp_remat", "expt_gtp_remat"])``; both + getters use ``check_initialized=False``, so an uninitialized GTP axis must yield None groups + rather than break construction of every non-GTP module. + """ + + def test_gtp_remat_pgs_are_none_and_do_not_raise(self, mocker): + """The exact call in the TE extension returns None groups, no assert, when GTP is off.""" + # Force the uninitialized GTP state deterministically (independent of suite ordering). + mocker.patch.object(parallel_state, "_GTP_WEIGHT_REMAT_GROUP", None) + mocker.patch.object(parallel_state, "_EXPERT_GTP_WEIGHT_REMAT_GROUP", None) + + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=["gtp_remat", "expt_gtp_remat"] + ) + + # Mirror the downstream selection in the TE extension; both branches are None, so + # _init_gtp_remat_context takes the no-op (GTP-inactive) path. + assert pg_collection.gtp_remat is None + assert pg_collection.expt_gtp_remat is None + for is_expert in (False, True): + gtp_remat_group = pg_collection.expt_gtp_remat if is_expert else pg_collection.gtp_remat + assert gtp_remat_group is None + + +# --------------------------------------------------------------------------- +# weight_prefetch=False: forward still produces correct output +# --------------------------------------------------------------------------- + + +def _worker_prefetch_disabled(rank, world_size, port): + torch.manual_seed(0) + in_f, out_f = 32, 64 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + gtp_module.update_gtp_config(weight_prefetch=False) + try: + l0 = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + l1 = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + + inp = torch.randn(4, in_f, dtype=dtype, device="cuda") + dist.broadcast(inp, src=0) + + # Single forward pass: builds chain and verifies output is correct + out = l0(inp, is_first_microbatch=True) + l1(inp, is_first_microbatch=True) + + # Chain should still be wired even with prefetch disabled + assert l0.weight.next_w is l1.weight + assert torch.isfinite(out).all(), "Non-finite output with prefetch disabled" + finally: + gtp_module.update_gtp_config(weight_prefetch=True) + + +class TestGTPPrefetchDisabled: + def test_forward_works_without_prefetch(self): + _requires_multi_gpu(4) + _run_distributed(_worker_prefetch_disabled, 4) + + +# --------------------------------------------------------------------------- +# fuse_wgrad_accumulation=True: wgrad is accumulated into main_grad +# --------------------------------------------------------------------------- + + +def _worker_fuse_wgrad(rank, world_size, port): + torch.manual_seed(0) + in_f, out_f = 32, 128 # out_f % (16*world_size)==0, no padding + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + layer = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype, fuse_wgrad_accumulation=True) + + # Allocate main_grad on the local shard shape + w = layer.weight + w.main_grad = torch.zeros(w.shape, dtype=dtype, device="cuda") + + inp = torch.randn(8, in_f, dtype=dtype, device="cuda", requires_grad=True) + dist.broadcast(inp, src=0) + + layer(inp, is_first_microbatch=True).sum().backward() + + # With fused accumulation, wgrad was added into main_grad + assert torch.any( + w.main_grad != 0 + ), "main_grad should have been updated by fused wgrad accumulation" + + +class TestFuseWgradAccumulation: + def test_wgrad_accumulated_into_main_grad(self): + _requires_multi_gpu(4) + _run_distributed(_worker_fuse_wgrad, 4) + + +# --------------------------------------------------------------------------- +# _grad_accum_hook is called after reduce-scatter +# --------------------------------------------------------------------------- + + +def _worker_main_grad_updated_after_bwd(rank, world_size, port): + """After backward, the wgrad RS path must have accumulated wgrad into main_grad.""" + torch.manual_seed(0) + in_f, out_f = 32, 64 + dtype = torch.bfloat16 + gtp_remat_group = dist.new_group(list(range(world_size))) + + layer = _make_gtp_linear(in_f, out_f, gtp_remat_group, dtype) + + # wgrad RS path always accumulates into main_grad; allocate before backward. + layer.weight.main_grad = torch.zeros(layer.weight.shape, dtype=dtype, device="cuda") + + inp = torch.randn(8, in_f, dtype=dtype, device="cuda", requires_grad=True) + dist.broadcast(inp, src=0) + layer(inp, is_first_microbatch=True).sum().backward() + + assert torch.any( + layer.weight.main_grad != 0 + ), "main_grad should have been updated after the reduce-scatter accumulation" + + +class TestGTPGradAccumHook: + def test_main_grad_updated_after_backward(self): + _requires_multi_gpu(4) + _run_distributed(_worker_main_grad_updated_after_bwd, 4) + + +# --------------------------------------------------------------------------- +# wait_async_comms(finalize_after_drain=True) inline-accumulation fallback +# --------------------------------------------------------------------------- + + +class TestWaitAsyncCommsFallback: + """Exercises the inline-accumulation fallback inside + ``wait_async_comms(finalize_after_drain=True)``: when a param is in + ``_inflight_comm_params`` (async AG was issued) but its ``_wgrad_rs_handle`` + is ``None`` (no async RS handle to drain), the inner + ``_wait_reduce_scatter`` call no-ops and the outer loop must inline the + accumulation itself (main_grad.add_ + ticket release + flag set). + + Production flows rarely hit this combination — chain-interior params have + both async AG and async RS, and chain-head sync RS doesn't enter + ``_inflight_comm_params`` via bwd AG. We construct the state by hand to + pin down the fallback's contract. + """ + + @staticmethod + def _make_inflight_param(main_grad_fill=0.0, already_finalized=False): + """Build a minimal GTPShardedParam wired for wait_async_comms testing.""" + dtype = torch.bfloat16 + p = GTPShardedParam(torch.zeros(8, 4, dtype=dtype, device="cuda")) + p.group = _FakeGroup() + p.expert_idx = None + p.pad_length = 0 + p.chain_id = gtp_module.GTPChain.UNGRAPHED.value + p._quantizer = None + p.is_routed_expert = False # ⇒ self._weights property returns [self] + p.main_grad = torch.full((8, 4), main_grad_fill, dtype=dtype, device="cuda") + p._prefetch_handle = None # _wait_param_gather is no-op + p._wgrad_rs_handle = None # _wait_reduce_scatter is no-op → fallback fires + p._cached_ag_stream = None + p._cached_rs_stream = None + p.ag_event = torch.cuda.Event(external=True) + p.rs_event = torch.cuda.Event(external=True) + p.rs_event.record() # so rs_event.wait() in fallback doesn't block + p._already_finalized = already_finalized + p.grad_added_to_main_grad = False + return p + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + def test_fallback_accumulates_when_no_rs_handle(self): + dtype = torch.bfloat16 + p = self._make_inflight_param(main_grad_fill=0.0) + + # Place a known wgrad in the cache for the fallback to read. + cache = gtp_module.get_global_GTP_cache() + p._rs_ticket = cache.reserve(p, dtype, fwd=False, reduce_scatter=True) + cache.get(p._rs_ticket).fill_(2.0) + + # Save + replace _inflight_comm_params so we don't trip over leftover + # params from earlier tests in the loop. + saved = set(gtp_module._inflight_comm_params) + gtp_module._inflight_comm_params.clear() + gtp_module._inflight_comm_params.add(p) + try: + gtp_module.wait_async_comms( + chain_id=p.chain_id, skip_rs=False, finalize_after_drain=True + ) + finally: + gtp_module._inflight_comm_params.clear() + gtp_module._inflight_comm_params.update(saved) + + torch.cuda.synchronize() + assert torch.all( + p.main_grad == 2.0 + ), f"main_grad should be 2.0 after fallback accumulation; got {p.main_grad}" + assert p._already_finalized is True, "_already_finalized must be set" + assert p.grad_added_to_main_grad is True, "grad_added_to_main_grad must be set" + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + def test_fallback_skipped_when_already_finalized(self): + """When _already_finalized=True, the fallback must NOT re-accumulate.""" + p = self._make_inflight_param(main_grad_fill=5.0, already_finalized=True) + # No _rs_ticket: if the fallback ran it would AttributeError on cache.get(None). + p._rs_ticket = None + + saved = set(gtp_module._inflight_comm_params) + gtp_module._inflight_comm_params.clear() + gtp_module._inflight_comm_params.add(p) + try: + gtp_module.wait_async_comms( + chain_id=p.chain_id, skip_rs=False, finalize_after_drain=True + ) + finally: + gtp_module._inflight_comm_params.clear() + gtp_module._inflight_comm_params.update(saved) + + torch.cuda.synchronize() + assert torch.all( + p.main_grad == 5.0 + ), "main_grad must be untouched when _already_finalized=True" + + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") + def test_fallback_skipped_for_pure_ag_param(self): + """Regression: cross-graph fwd-AG prefetch in flight + finalize_after_drain=True. + + A param can be in _inflight_comm_params because of an outstanding async + all-gather (e.g. a cross-graph forward prefetch reaching the + bwd→optimizer boundary). No reduce-scatter was ever issued for that + param, so _rs_ticket is None on every weight. Previously the fallback + called cache.get(None) and crashed with KeyError; the guard now skips + the inline accumulation entirely when no weight has an RS ticket. + """ + p = self._make_inflight_param(main_grad_fill=7.0) + # Critical: simulates a pure-AG prefetch — no RS ever issued, ticket is None. + p._rs_ticket = None + + saved = set(gtp_module._inflight_comm_params) + gtp_module._inflight_comm_params.clear() + gtp_module._inflight_comm_params.add(p) + try: + # Must NOT raise KeyError(None) from cache.get(None). + gtp_module.wait_async_comms( + chain_id=p.chain_id, skip_rs=False, finalize_after_drain=True + ) + finally: + gtp_module._inflight_comm_params.clear() + gtp_module._inflight_comm_params.update(saved) + + torch.cuda.synchronize() + assert torch.all( + p.main_grad == 7.0 + ), "main_grad must be untouched for a pure-AG param (no wgrad to accumulate)" + assert ( + p._already_finalized is False + ), "_already_finalized must stay False — no finalize happened for a pure-AG param" + + +# --------------------------------------------------------------------------- +# GTP_remat DDP bucket alignment: distributed optimizer bucket-end assertion +# --------------------------------------------------------------------------- + + +def _worker_gtp_ddp_bucket_alignment(rank, world_size, port): + """GTP param buffers in DDP must use padded bucket layout with use_distributed_optimizer=True. + + Bug: DDP used param_layout=None for GTP buffers, falling through to + _compute_default_per_buffer_param_layout, which packs params without padding bucket ends. + The distributed optimizer requires every bucket end to be divisible by + intra_dp_cp_group.size() (asserted at param_and_grad_buffer.py:1427). + + Trigger: + GTP_remat_size=2, DP=4 → intra_dp_cp_group.size()=2 + pad_for_alignment=0, weight [out=2,in=3] → GTP shard=[1,3]=3 elements (odd) + Two GTP params: total=6, 6%2==0 (total check passes); bucket_size=3 forces + bucket-0 to contain only the first param, end=3, 3%2≠0 → AssertionError + """ + from megatron.core import parallel_state as ps + from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig + from megatron.core.transformer.transformer_config import TransformerConfig + + # The module fixture initialized model_parallel without GTP_remat; re-init with GTP_remat=2. + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + + orig_pad = gtp_module.GTP_CONFIG.pad_for_alignment + gtp_module.GTP_CONFIG.pad_for_alignment = 0 + try: + gtp_remat_group = ps.get_gtp_weight_remat_group() + + class _TwoLayerModel(torch.nn.Module): + def __init__(self): + super().__init__() + self.fc0 = te.Linear(3, 2, bias=False, device="cuda") + self.fc1 = te.Linear(3, 2, bias=False, device="cuda") + + model = _TwoLayerModel() + wrap_module_params_gtp(model.fc0, ["weight"], gtp_remat_group) + wrap_module_params_gtp(model.fc1, ["weight"], gtp_remat_group) + + config = TransformerConfig( + num_attention_heads=1, num_layers=1, hidden_size=4, tensor_model_parallel_size=1 + ) + ddp_config = DistributedDataParallelConfig( + use_distributed_optimizer=True, overlap_grad_reduce=True, bucket_size=3 + ) + + # Without the fix this raises AssertionError at param_and_grad_buffer.py:1427: + # assert end_index % self.data_parallel_world_size == 0 + DistributedDataParallel(config, ddp_config, model) + finally: + gtp_module.GTP_CONFIG.pad_for_alignment = orig_pad + ps.destroy_model_parallel() + ps.initialize_model_parallel() # restore default for remaining tests + + +def _worker_regular_buffer_padded_when_gtp_params_present(rank, world_size, port): + """Regular (non-GTP) param buffers in DDP must also use padded layout when GTP is active. + + Bug: when gtp_params is non-empty, full_param_layout.layouts contains stale GTP entries + that don't belong to the regular buffer, causing KeyErrors in DistOpt's param map. + DDP avoided this by forcing param_layout=None for regular buffers, but that falls through + to _compute_default_per_buffer_param_layout, which produces unpadded bucket ends, again + violating param_and_grad_buffer.py:1427 (end_index % data_parallel_world_size == 0). + + Trigger: + GTP_remat_size=2, DP=4 → intra_dp_cp_group.size()=4 + (regular params reduce over the full DP group) + bias=True → each bias has 2 elements (not divisible by 4) + Two layers: total regular numel=4, 4%4==0 (total check passes); bucket_size=2 forces + bucket-0 to contain only the first bias, end=2, 2%4≠0 → AssertionError + """ + from megatron.core import parallel_state as ps + from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig + from megatron.core.transformer.transformer_config import TransformerConfig + + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + + orig_pad = gtp_module.GTP_CONFIG.pad_for_alignment + gtp_module.GTP_CONFIG.pad_for_alignment = 0 + try: + gtp_remat_group = ps.get_gtp_weight_remat_group() + + class _TwoLayerModelWithBias(torch.nn.Module): + def __init__(self): + super().__init__() + # bias=True: weight → GTPShardedParam (gtp_buffer), bias → regular param + self.fc0 = te.Linear(3, 2, bias=True, device="cuda") + self.fc1 = te.Linear(3, 2, bias=True, device="cuda") + + model = _TwoLayerModelWithBias() + wrap_module_params_gtp(model.fc0, ["weight"], gtp_remat_group) + wrap_module_params_gtp(model.fc1, ["weight"], gtp_remat_group) + + config = TransformerConfig( + num_attention_heads=1, num_layers=1, hidden_size=4, tensor_model_parallel_size=1 + ) + # bucket_size=2: each 2-element bias fills one bucket in the regular buffer. + # Without the fix: regular buffer uses param_layout=None → bucket-0 ends at 2, + # 2 % intra_dp_cp_group.size()(=4) != 0 → AssertionError at line 1427. + ddp_config = DistributedDataParallelConfig( + use_distributed_optimizer=True, overlap_grad_reduce=True, bucket_size=2 + ) + + DistributedDataParallel(config, ddp_config, model) + finally: + gtp_module.GTP_CONFIG.pad_for_alignment = orig_pad + ps.destroy_model_parallel() + ps.initialize_model_parallel() + + +class TestGTPDDPBucketAlignment: + def test_gtp_buffers_use_padded_layout_with_distributed_optimizer(self): + """GTP buffer bucket ends must be padded to intra_dp_cp_group.size().""" + _requires_multi_gpu(4) + _run_distributed(_worker_gtp_ddp_bucket_alignment, 4) + + def test_regular_buffers_use_padded_layout_when_gtp_params_present(self): + """Regular buf bucket ends must be padded even when gtp_params forces layoutrecompute.""" + _requires_multi_gpu(4) + _run_distributed(_worker_regular_buffer_padded_when_gtp_params_present, 4) + + +# --------------------------------------------------------------------------- +# GTP_remat DDP grad-ready wiring: register_grad_ready must fire AFTER the wgrad add +# --------------------------------------------------------------------------- + + +def _worker_gtp_ddp_grad_ready_wiring(rank, world_size, port): + """GTP params must drive DDP grad-ready from GTP's manual hook, not autograd. + + GTP defers the main_grad accumulation to a later backward node, so autograd's AccumulateGrad can + fire register_grad_ready before the grad lands and dispatch the bucket reduce-scatter on stale + grad_data (corrupts reduce_scatter_with_fp32_accumulation). The fix routes grad-ready through + register_grad_accum_hook (fired after the add) and skips the autograd hook. This pins that + wiring: every GTP weight has _grad_accum_hook set and none falls through to the autograd list. + """ + from megatron.core import parallel_state as ps + from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig + from megatron.core.transformer.transformer_config import TransformerConfig + + # The module fixture initialized model_parallel without GTP_remat; re-init with GTP_remat=2. + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + try: + gtp_remat_group = ps.get_gtp_weight_remat_group() + + class _TwoLayerModel(torch.nn.Module): + def __init__(self): + super().__init__() + # bias=False -> all params are GTP_remat weights, so grad_accs must end up empty. + self.fc0 = te.Linear(64, 128, bias=False, device="cuda") + self.fc1 = te.Linear(64, 128, bias=False, device="cuda") + + model = _TwoLayerModel() + wrap_module_params_gtp(model.fc0, ["weight"], gtp_remat_group) + wrap_module_params_gtp(model.fc1, ["weight"], gtp_remat_group) + + config = TransformerConfig( + num_attention_heads=1, num_layers=1, hidden_size=4, tensor_model_parallel_size=1 + ) + ddp_config = DistributedDataParallelConfig( + use_distributed_optimizer=True, overlap_grad_reduce=True + ) + ddp_model = DistributedDataParallel(config, ddp_config, model) + + for name, w in [("fc0", model.fc0.weight), ("fc1", model.fc1.weight)]: + assert isinstance(w, GTPShardedParam), f"{name}.weight should be a GTP param" + # Manual hook set -> grad-ready fires after the add; None -> early autograd path (bug). + assert ( + getattr(w, "_grad_accum_hook", None) is not None + ), f"{name}.weight must have _grad_accum_hook set (manual grad-ready, not autograd)" + + # bias=False -> all params are GTP_remat -> none took the autograd path. + assert len(ddp_model.grad_accs) == 0, ( + "GTP params must not register an autograd AccumulateGrad hook " + f"(grad_accs has {len(ddp_model.grad_accs)} entries)" + ) + finally: + ps.destroy_model_parallel() + ps.initialize_model_parallel() # restore default for remaining tests + + +class TestGTPDDPGradReadyWiring: + def test_gtp_params_use_manual_grad_ready_hook(self): + """GTP params route DDP grad-ready through register_grad_accum_hook, not autograd.""" + _requires_multi_gpu(4) + _run_distributed(_worker_gtp_ddp_grad_ready_wiring, 4) diff --git a/tests/unit_tests/generalized_tensor_parallel/test_gtp_cudagraph_grad.py b/tests/unit_tests/generalized_tensor_parallel/test_gtp_cudagraph_grad.py new file mode 100644 index 00000000000..d7c2f4e53e4 --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_gtp_cudagraph_grad.py @@ -0,0 +1,95 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Regression test for the GTP + CUDA-graph capture-step grad-norm bug. + +Bug: create_cudagraphs() runs after finalize_model_grads, so main_grad already holds the finalized +(reduced + per-token-scaled) grads. create_fwd_graph then runs an eager warmup backward (graph +capture only records ops, it doesn't run them), and that eager backward executes GTP's wgrad +main_grad.add_ -- including the cascade add into a param's cross-graph ``next_w`` (in another +module, via a stale RS ticket) -- clobbering the finalized grads and spiking the step's grad norm. + +Fix: create_fwd_graph snapshots the grads its warmup touches via ``_backup_grads_before_capture`` +and restores them after. This test exercises that helper pair directly: the module's own params +and their cross-graph ``next_w`` must survive a simulated warmup clobber. +""" + +import pytest +import torch + +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP +from megatron.core.transformer.cuda_graphs import ( + _backup_grads_before_capture, + _restore_grads_after_capture, +) + +if not HAVE_GTP: + pytest.skip("GTP requires TE with hook registry", allow_module_level=True) + + +def _gtp_param(value: float, numel: int = 8) -> torch.nn.Parameter: + """A param with a finalized (reduced + scaled) main_grad, flagged as a GTP weight.""" + p = torch.nn.Parameter(torch.zeros(numel, device="cuda")) + p.is_gtp_weight_remat = True + p.main_grad = torch.full((numel,), value, device="cuda") + return p + + +class _Mod(torch.nn.Module): + def __init__(self, weight: torch.nn.Parameter): + super().__init__() + self.weight = weight + + +class _StubRunner: + """The ``base_module`` and ``gtp_remat`` attrs that ``_backup_grads_before_capture`` reads.""" + + def __init__(self, base_module: torch.nn.Module, gtp_remat: bool = True): + self.base_module = base_module + self.gtp_remat = gtp_remat + + +class TestGTPCaptureGradSnapshot: + def test_preserves_own_and_cross_graph_next_w(self): + """Snapshot/restore must keep both the module's own grad and its cross-graph next_w grad + (in another module) intact across a capture that clobbers them.""" + own = _gtp_param(0.0125) + cross = _gtp_param(0.02) # next_w lives in a different module/graph + own.next_w = cross + runner = _StubRunner(_Mod(own)) + + backup = _backup_grads_before_capture(runner) + own.main_grad.add_(410.0) # simulate the capture-time main_grad.add_ clobber + cross.main_grad.add_(99.0) + _restore_grads_after_capture(backup) + + torch.testing.assert_close(own.main_grad, torch.full((8,), 0.0125, device="cuda")) + torch.testing.assert_close(cross.main_grad, torch.full((8,), 0.02, device="cuda")) + + def test_routed_expert_next_w_via_weight_list(self): + """A routed-expert next_w exposes its shards via ``weight_list`` (read directly, since the + ``_weights`` property raises on non-leaders before capture).""" + own = _gtp_param(0.0125) + shard0, shard1 = _gtp_param(0.03), _gtp_param(0.04) + routed = torch.nn.Parameter(torch.zeros(8, device="cuda")) # leader wrapper (no own grad) + routed.is_routed_expert = True + routed.weight_list = [shard0, shard1] + own.next_w = routed + runner = _StubRunner(_Mod(own)) + + backup = _backup_grads_before_capture(runner) + shard0.main_grad.add_(50.0) + shard1.main_grad.add_(60.0) + _restore_grads_after_capture(backup) + + torch.testing.assert_close(shard0.main_grad, torch.full((8,), 0.03, device="cuda")) + torch.testing.assert_close(shard1.main_grad, torch.full((8,), 0.04, device="cuda")) + + def test_non_gtp_backs_up_own_params_only(self): + """Non-GTP runner: own params are snapshotted, but the GTP cross-graph next_w walk is + skipped (the bwd capture doesn't touch main_grad on the non-GTP path).""" + own = _gtp_param(0.0125) + cross = _gtp_param(0.02) + own.next_w = cross + backup = _backup_grads_before_capture(_StubRunner(_Mod(own), gtp_remat=False)) + assert id(own) in backup + assert id(cross) not in backup diff --git a/tests/unit_tests/generalized_tensor_parallel/test_gtp_dcp.py b/tests/unit_tests/generalized_tensor_parallel/test_gtp_dcp.py new file mode 100644 index 00000000000..da42a086a91 --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_gtp_dcp.py @@ -0,0 +1,1089 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Unit tests for GTP_remat + distributed checkpointing. + +Verifies that ``make_sharded_tensors_for_checkpoint_with_gtp_remat`` emits +ShardedTensor offsets that correctly encode TP × GTP_remat sharding, and that +the helper is a no-op (delegates to vanilla) when no ``GTPShardedParam`` +is present in the input state_dict. + +""" + +import pytest +import torch +import torch.distributed as dist + +from megatron.core import parallel_state as ps +from megatron.core.dist_checkpointing import ShardedTensor +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + +if not HAVE_GTP: + pytest.skip("GTP requires TE with hook registry", allow_module_level=True) + +import transformer_engine.pytorch as te # noqa: E402 +from transformer_engine.common.recipe import MXFP8BlockScaling # noqa: E402 +from transformer_engine.pytorch import fp8_autocast, fp8_model_init # noqa: E402 + +from megatron.core.dist_checkpointing.mapping import ( # noqa: E402 + ShardedObject, + ShardedTensorFactory, + is_main_replica, +) +from megatron.core.extensions.transformer_engine import ( # noqa: E402 + TELayerNormColumnParallelLinear, + TERowParallelLinear, +) +from megatron.core.fp8_utils import is_float8tensor # noqa: E402 +from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add # noqa: E402 +from megatron.core.process_groups_config import ProcessGroupCollection # noqa: E402 +from megatron.core.ssm.mamba_layer import MambaLayer, MambaLayerSubmodules # noqa: E402 +from megatron.core.ssm.mamba_mixer import MambaMixer, MambaMixerSubmodules # noqa: E402 +from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( # noqa: E402 + GTP_CONFIG, + GTPShardedParam, + make_sharded_tensors_for_checkpoint_with_gtp_remat, + update_gtp_config, + wrap_module_params_gtp, +) +from megatron.core.tensor_parallel.gtp_api import ( # noqa: E402 + attach_gtp_to_presharded_module, + dequantize_gtp_native_fp8, + gtp_native_fp8_load_context, + gtp_remat_shard_dim0, +) +from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed # noqa: E402 +from megatron.core.transformer.spec_utils import ModuleSpec # noqa: E402 +from megatron.core.transformer.transformer_config import TransformerConfig # noqa: E402 +from megatron.core.transformer.utils import make_sharded_tensors_for_checkpoint # noqa: E402 +from megatron.core.utils import get_pg_size, make_tp_sharded_tensor_for_checkpoint # noqa: E402 +from tests.unit_tests.generalized_tensor_parallel.gtp_test_utils import ( # noqa: E402,F401 + _requires_mxfp8, + _torchrun_dist_init, +) + + +@pytest.fixture(autouse=True) +def _no_pad_alignment(): + """Disable GTP_remat padding for the duration of each test so local shard sizes + are exactly ``per_tp_out / gtp_remat_size`` and the test math stays simple. + DCP semantics with padding are exercised by the integration tests. + """ + orig = GTP_CONFIG.pad_for_alignment + update_gtp_config(pad_for_alignment=0) + yield + update_gtp_config(pad_for_alignment=orig) + + +def _require_world_size(n): + if dist.get_world_size() != n: + pytest.skip( + f"Requires world_size={n}, got {dist.get_world_size()} " + f"(launch with torchrun --nproc-per-node={n})" + ) + + +# Many workers need the same TP/GTP subgroups. Memoize by rank-set so the process holds a +# handful of communicators instead of re-creating (and leaking) one per worker. +_GROUP_CACHE = {} + + +def _cached_new_group(ranks): + """Memoized ``dist.new_group`` keyed by rank-set (see note above).""" + key = tuple(ranks) + if key not in _GROUP_CACHE: + _GROUP_CACHE[key] = dist.new_group(list(ranks)) + return _GROUP_CACHE[key] + + +@pytest.fixture(scope="module", autouse=True) +def _precreate_subgroups(_torchrun_dist_init): + """Pre-create the shared TP/GTP subgroups once, on all ranks, in a fixed order. + + ``dist.new_group`` is a world-collective (all ranks must call it in the same order); the + per-member ``new_group([0,1]) if rank in (0,1) else ...`` idiom collides disjoint groups on + the call-order tag and hangs NCCL. Pre-creating makes every later ``_cached_new_group`` a hit. + """ + if dist.is_initialized() and dist.get_world_size() == 4: + for ranks in ([0, 1], [2, 3], [0, 2], [1, 3], [0, 1, 2, 3]): + _cached_new_group(ranks) + yield + + +def _make_gtp_shard(out_features, in_features, gtp_remat_group, dtype=torch.bfloat16): + """Build a small GTPShardedParam by wrapping a one-param dummy module.""" + + class _Dummy(torch.nn.Module): + def __init__(self): + super().__init__() + self.weight = torch.nn.Parameter( + torch.arange(out_features * in_features, dtype=dtype, device="cuda").reshape( + out_features, in_features + ) + ) + + mod = _Dummy() + wrap_module_params_gtp(mod, ["weight"], gtp_remat_group) + return mod.weight # now a GTPShardedParam + + +def _make_native_fp8_gtp_shard(per_tp_out, in_f, gtp_remat_group, recipe): + """Build a native-FP8 GTP weight the production way (extensions/transformer_engine.py): + pass the pre-sharded out_features into a stock ``fp8_model_init`` ``te.Linear`` so TE inits a + native MXFP8 shard, attach the GTP surface post-init, then run one FP8 forward to populate the + rowwise/columnwise FP8 data. Returns the reclassed ``GTP_`` weight.""" + shard_out, pad = gtp_remat_shard_dim0(per_tp_out, gtp_remat_group) + with fp8_model_init(enabled=True, recipe=recipe): + lin = te.Linear(in_f, shard_out, bias=False, params_dtype=torch.bfloat16, device="cuda") + lin.gtp_remat_size = gtp_remat_group.size() + attach_gtp_to_presharded_module(lin, gtp_remat_group, pad) + with fp8_autocast(enabled=True, fp8_recipe=recipe): + _ = lin(torch.randn(32, in_f, dtype=torch.bfloat16, device="cuda")) + return lin.weight + + +def _worker_native_fp8_dcp_save(rank, world_size, port): + """Native-FP8 GTP weight: DCP save must emit a dequantized BF16 ShardedTensor with the full + (TP x GTP_remat) global shape and correct composite axis-0 offset -- not raw FP8 bytes under + a fake BF16 dtype (a55b save-crash guard: recognition gates / TE tex.dequantize miss the + native-FP8 GTP_ subclass; dequantize_gtp_native_fp8 restores the base class). + """ + _requires_mxfp8() + + # TP=2, GTP_remat=2 (4 ranks). MXFP8 needs dims % 32, so use fp8-valid sizes. + gtp_remat_group = _cached_new_group([0, 1]) if rank in (0, 1) else _cached_new_group([2, 3]) + tp_group = _cached_new_group([0, 2]) if rank in (0, 2) else _cached_new_group([1, 3]) + full_out, in_f = 128, 128 + tp_size, gtp_remat_size = 2, 2 + per_tp_out = full_out // tp_size # 64 + per_shard_out = per_tp_out // gtp_remat_size # 32 (== MXFP8 block size) + + recipe = MXFP8BlockScaling() + w = _make_native_fp8_gtp_shard(per_tp_out, in_f, gtp_remat_group, recipe) + # The live weight is a native FP8 GTP param (a QuantizedTensor subclass), sharded. + assert is_float8tensor(w), "weight should be a native FP8 tensor" + assert getattr(w, "is_gtp_weight_remat", False), "GTP surface missing" + assert type(w).__name__.startswith("GTP_"), type(w).__name__ + assert tuple(w.shape) == (per_shard_out, in_f) + + sharded = make_sharded_tensors_for_checkpoint_with_gtp_remat( + {"weight": w}, + prefix="", + tensor_parallel_layers_axis_map={"weight": 0}, + sharded_offsets=(), + tp_group=tp_group, + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + st = sharded["weight"] + assert isinstance(st, ShardedTensor), type(st) + + # Saved data must be dequantized BF16 — not raw FP8 bytes under a fake dtype. + assert st.data.dtype == torch.bfloat16, f"expected bf16 saved data, got {st.data.dtype}" + assert not is_float8tensor(st.data), "checkpoint data must be dequantized, not FP8" + assert tuple(st.data.shape) == (per_shard_out, in_f) + + # Full (TP x GTP) global shape + composite axis-0 offset (not sharded-as-full). + assert st.global_shape[0] == full_out, (st.global_shape, full_out) + tp_rank, gtp_rank = rank // 2, rank % 2 + assert st.global_offset[0] == (tp_rank * gtp_remat_size + gtp_rank) * per_shard_out, ( + rank, + st.global_offset, + ) + + # The live param must be untouched by the save (class restored, still native FP8). + assert is_float8tensor(w) and type(w).__name__.startswith( + "GTP_" + ), "dequantize must not mutate the live param's class" + + +def _worker_native_fp8_dcp_load_copy(rank, world_size, port): + """Copying a BF16 checkpoint value back into a live native-FP8 GTP weight must go through + ``gtp_native_fp8_load_context`` (a55b load-crash guard: TE's exact-class MXFP8 check rejects + the dynamic ``GTP_`` subclass). Assert the raw copy raises but succeeds under the + context, and the reclassed weight dequantizes to the loaded values. + """ + _requires_mxfp8() + + # This test exercises a single-rank concern (the __class__ swap during copy_), so use the + # default WORLD group as the gtp_remat_group rather than dist.new_group subgroups — the + # latter's secondary NCCL socket bootstrap is flaky on some multi-node allocations and would + # mask the fp8 copy behavior under test. + gtp_remat_group = dist.group.WORLD + per_tp_out, in_f = 128, 128 # MXFP8 needs dims % 32; shard = 128/world(4) = 32 + recipe = MXFP8BlockScaling() + + shard_out, pad = gtp_remat_shard_dim0(per_tp_out, gtp_remat_group) + with fp8_model_init(enabled=True, recipe=recipe): + lin = te.Linear(in_f, shard_out, bias=False, params_dtype=torch.bfloat16, device="cuda") + lin.gtp_remat_size = gtp_remat_group.size() + attach_gtp_to_presharded_module(lin, gtp_remat_group, pad) + with fp8_autocast(enabled=True, fp8_recipe=recipe): + _ = lin(torch.randn(32, in_f, dtype=torch.bfloat16, device="cuda")) + + assert is_float8tensor(lin.weight) and type(lin.weight).__name__.startswith("GTP_") + + # The dequantized BF16 payload a DCP load would hand back for this shard. + target_bf16 = torch.randn(shard_out, in_f, dtype=torch.bfloat16, device="cuda") + + # (1) Without the context, copy_ into the subclass raises in TE's C++ quantizer. + # Mirror production's _load_from_state_dict, which copies under no_grad. + raised = False + try: + with torch.no_grad(): + lin.weight.copy_(target_bf16) + except Exception as e: # noqa: BLE001 + raised = True + assert "MXFP8" in str(e) or "IsMXFP8Tensor" in str(e), str(e) + assert raised, "copy_ into GTP_ unexpectedly succeeded without the load context" + + # (2) Under the context the copy succeeds; the reclassed weight holds the loaded values. + with torch.no_grad(), gtp_native_fp8_load_context(lin): + lin.weight.copy_(target_bf16) + assert is_float8tensor(lin.weight) and type(lin.weight).__name__.startswith( + "GTP_" + ), "load context must reclass back to the GTP subclass" + loaded = dequantize_gtp_native_fp8(lin.weight) + # MXFP8 round-trip is lossy; check it tracks the target (not the pre-copy garbage). + rel = (loaded - target_bf16).abs().max() / target_bf16.abs().max().clamp_min(1e-6) + assert rel < 0.2, f"loaded weight does not match checkpoint values (max rel {rel:.3f})" + + +def _worker_helper_offsets_tp_eq_gtp_axis(rank, world_size, port): + """TP=2, GTP_remat=2 (4 ranks total). Weight is GTPShardedParam. + + Production flow: Mcore TE constructs the Linear with already-TP-sliced + out_features (i.e. full / tp_size). GTP_remat then slices that further by + gtp_remat_size. We mimic that by starting with a per-TP-rank tensor of size + ``full // tp_size`` and letting wrap_module_params_gtp slice it. + """ + gtp_remat_group = _cached_new_group([0, 1]) if rank in (0, 1) else _cached_new_group([2, 3]) + tp_group = _cached_new_group([0, 2]) if rank in (0, 2) else _cached_new_group([1, 3]) + + full_out_features = 8 + tp_size, gtp_remat_size = 2, 2 + per_tp_out = full_out_features // tp_size # 4 + per_shard_out = per_tp_out // gtp_remat_size # 2 + in_features = 4 + + weight = _make_gtp_shard(per_tp_out, in_features, gtp_remat_group) + assert weight.shape == (per_shard_out, in_features), ( + f"rank={rank} local shard shape {tuple(weight.shape)} != " + f"({per_shard_out}, {in_features})" + ) + + sharded = make_sharded_tensors_for_checkpoint_with_gtp_remat( + {"weight": weight}, + prefix="", + tensor_parallel_layers_axis_map={"weight": 0}, + sharded_offsets=(), + tp_group=tp_group, + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + st = sharded["weight"] + assert isinstance(st, ShardedTensor), f"Expected ShardedTensor, got {type(st)}" + + # Composite offset: (axis=0, tp_rank*gtp_remat_size+gtp_rank, tp_size*gtp_remat_size) + # rank → (tp_rank, gtp_rank): 0→(0,0), 1→(0,1), 2→(1,0), 3→(1,1) + tp_rank = rank // 2 + gtp_rank = rank % 2 + expected_offset = (tp_rank * gtp_remat_size + gtp_rank) * per_shard_out + assert ( + st.global_offset[0] == expected_offset + ), f"rank={rank} expected axis-0 offset {expected_offset}, got {st.global_offset[0]}" + assert ( + st.global_shape[0] == full_out_features + ), f"rank={rank} expected global axis-0 size {full_out_features}, got {st.global_shape[0]}" + + +def _worker_helper_offsets_tp_neq_gtp_axis(rank, world_size, port): + """Row-parallel: TP=2 shards axis 1, GTP_remat=2 shards axis 0. + + Per-TP-rank tensor: (full_out, full_in/tp_size). GTP_remat further shards + axis 0 to (full_out/gtp_remat_size, full_in/tp_size). + """ + gtp_remat_group = _cached_new_group([0, 1]) if rank in (0, 1) else _cached_new_group([2, 3]) + tp_group = _cached_new_group([0, 2]) if rank in (0, 2) else _cached_new_group([1, 3]) + + full_out, full_in = 8, 4 + tp_size, gtp_remat_size = 2, 2 + per_tp_in = full_in // tp_size # 2 + per_shard_out = full_out // gtp_remat_size # 4 + + weight = _make_gtp_shard(full_out, per_tp_in, gtp_remat_group) + assert weight.shape == (per_shard_out, per_tp_in) + + sharded = make_sharded_tensors_for_checkpoint_with_gtp_remat( + {"weight": weight}, + prefix="", + tensor_parallel_layers_axis_map={"weight": 1}, # row-parallel + sharded_offsets=(), + tp_group=tp_group, + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + st = sharded["weight"] + tp_rank = rank // 2 + gtp_rank = rank % 2 + assert ( + st.global_offset[0] == gtp_rank * per_shard_out + ), f"rank={rank} axis-0 offset wrong: {st.global_offset[0]}" + assert ( + st.global_offset[1] == tp_rank * per_tp_in + ), f"rank={rank} axis-1 offset wrong: {st.global_offset[1]}" + assert st.global_shape == ( + full_out, + full_in, + ), f"rank={rank} global shape {st.global_shape} != ({full_out}, {full_in})" + + +def _worker_helper_no_op_no_gtp_remat(rank, world_size, port): + """Helper must delegate to vanilla when state_dict has no GTPShardedParam. + + Per-TP-rank shape under column-parallel TP=2: (full_out//tp_size, in). + """ + tp_group = _cached_new_group([0, 1]) if rank in (0, 1) else _cached_new_group([2, 3]) + + full_out, in_features, tp_size = 8, 4, 2 + per_tp_out = full_out // tp_size + + plain = torch.nn.Parameter( + torch.zeros(per_tp_out, in_features, dtype=torch.bfloat16, device="cuda") + ) + bias = torch.nn.Parameter(torch.zeros(per_tp_out, dtype=torch.bfloat16, device="cuda")) + + sharded = make_sharded_tensors_for_checkpoint_with_gtp_remat( + {"weight": plain, "bias": bias}, + prefix="", + tensor_parallel_layers_axis_map={"weight": 0, "bias": 0}, + sharded_offsets=(), + tp_group=tp_group, + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + # tp_group is [0,1] for ranks 0,1 and [2,3] for ranks 2,3 here — local tp_rank = rank % 2 + tp_rank = rank % 2 + assert sharded["weight"].global_offset[0] == tp_rank * per_tp_out, ( + f"rank={rank} fallback path produced wrong offset for weight: " + f"{sharded['weight'].global_offset[0]}" + ) + assert sharded["weight"].global_shape == (full_out, in_features) + + +def _worker_helper_padded_inproj_no_pad_case(rank, world_size, port): + """``in_proj.weight`` shape modeled after the production case (z|x|B|C|dt + concat along dim 0). With GTP_remat=4 and these dim-0 sizes the alignment + constraint ``dim0 % (gtp_remat_size * pad_for_alignment) == 0`` is satisfied — + *no* padding fires. Verify the helper emits the expected offsets. + """ + update_gtp_config(pad_for_alignment=16) + # dim0 = 512+512+64+64+8 = 1160 → 1160 % (4*16=64) = 8 ⇒ NOT aligned. + # Pick sizes that ARE aligned to 64 to exercise the no-pad path: + dim0 = 1152 # = 18 * 64; alignment-clean for gtp_remat_size=4, pad=16 + in_features = 4 + + # All 4 ranks form a single GTP_remat group. + gtp_remat_group = _cached_new_group(list(range(world_size))) + weight = _make_gtp_shard(dim0, in_features, gtp_remat_group) + + # No padding ⇒ local shape is exactly dim0 / 4 = 288 + expected_local = dim0 // 4 + assert weight.shape == (expected_local, in_features), ( + f"rank={rank}: padding should NOT have fired (dim0 aligned); " + f"got local shape {tuple(weight.shape)}, expected ({expected_local}, {in_features})" + ) + assert getattr(weight, "pad_length", 0) == 0 + + sharded = make_sharded_tensors_for_checkpoint_with_gtp_remat( + {"weight": weight}, + prefix="", + tensor_parallel_layers_axis_map={"weight": 0}, + sharded_offsets=(), + tp_group=_cached_new_group([rank]), # trivial 1-rank TP group + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + st = sharded["weight"] + assert ( + st.global_shape[0] == dim0 + ), f"rank={rank} no-pad case: global_shape[0] {st.global_shape[0]} != {dim0}" + assert st.global_offset[0] == rank * expected_local + + +def _worker_helper_padded_inproj_pad_case(rank, world_size, port): + """in_proj with a dim-0 size needing GTP_remat padding (dim0=1160, gtp_remat_size=4, + pad_for_alignment=16 -> 56 pad rows -> padded 1216, per-rank shard 304). Pins that the + padded global shape round-trips when save_gtp_remat_size == load_gtp_remat_size. + """ + update_gtp_config(pad_for_alignment=16) + dim0_unpadded = 1160 # z(512) + x(512) + B(64) + C(64) + dt(8) + in_features = 4 + gtp_remat_size = world_size + alignment_block = 16 * gtp_remat_size # = 64 + pad = (alignment_block - dim0_unpadded % alignment_block) % alignment_block + dim0_padded = dim0_unpadded + pad + per_shard = dim0_padded // gtp_remat_size + + gtp_remat_group = _cached_new_group(list(range(world_size))) + weight = _make_gtp_shard(dim0_unpadded, in_features, gtp_remat_group) + + assert weight.shape == ( + per_shard, + in_features, + ), f"rank={rank}: post-pad shard shape {tuple(weight.shape)} != ({per_shard}, {in_features})" + # Only rank-3 (the last GTP_remat rank) carries the trailing pad rows; all ranks + # report the same pad_length (an invariant set by _gtp_slice_one_param). + assert ( + getattr(weight, "pad_length", 0) == pad + ), f"rank={rank}: pad_length {getattr(weight, 'pad_length', 0)} != {pad}" + + sharded = make_sharded_tensors_for_checkpoint_with_gtp_remat( + {"weight": weight}, + prefix="", + tensor_parallel_layers_axis_map={"weight": 0}, + sharded_offsets=(), + tp_group=_cached_new_group([rank]), + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + st = sharded["weight"] + # Helper saves the padded global. ``allow_shape_mismatch=True`` is what + # makes the saved tensor portable to a different load-time GTP_remat topology + # (different alignment choice yields a different padded size). + assert ( + st.global_shape[0] == dim0_padded + ), f"rank={rank} pad case: global_shape[0] {st.global_shape[0]} != {dim0_padded}" + assert st.global_offset[0] == rank * per_shard + assert st.allow_shape_mismatch is True, ( + f"rank={rank} pad case: allow_shape_mismatch must be True when GTP_remat padding fires; " + f"otherwise the ckpt cannot be loaded at a different GTP_remat topology." + ) + + +def _worker_helper_cross_topology_reshard_metadata(rank, world_size, port): + """Pin the cross-topology reshard contract via ShardedTensor metadata. + + We can't run a real DCP save/load against itself within a single torchrun + (need separate worlds), but we can verify the saved ShardedTensor carries + everything DCP needs to do the reshard: ``allow_shape_mismatch=True`` and + a global_shape large enough to cover any compatible load-side topology + (≥ unpadded original). + """ + update_gtp_config(pad_for_alignment=16) + dim0_unpadded = 1160 + in_features = 4 + gtp_remat_size = world_size + alignment_block = 16 * gtp_remat_size # 64 + dim0_padded = ( + dim0_unpadded + (alignment_block - dim0_unpadded % alignment_block) % alignment_block + ) + per_shard = dim0_padded // gtp_remat_size + + gtp_remat_group = _cached_new_group(list(range(world_size))) + weight = _make_gtp_shard(dim0_unpadded, in_features, gtp_remat_group) + + sharded = make_sharded_tensors_for_checkpoint_with_gtp_remat( + {"weight": weight}, + prefix="", + tensor_parallel_layers_axis_map={"weight": 0}, + sharded_offsets=(), + tp_group=_cached_new_group([rank]), + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + st = sharded["weight"] + # 1. The saved global covers >= unpadded original size. + assert st.global_shape[0] >= dim0_unpadded, ( + f"rank={rank} saved global_shape ({st.global_shape[0]}) < unpadded ({dim0_unpadded}); " + f"would lose valid data on cross-topology reshard." + ) + # 2. ``allow_shape_mismatch=True`` lets DCP tolerate that the load-side + # padded size may differ. + assert st.allow_shape_mismatch is True + # 3. Each rank's offset+local_shape covers a contiguous slice of the + # padded global; together the ranks cover [0, padded_global). + assert st.global_offset[0] + st.local_shape[0] <= st.global_shape[0] + assert st.global_offset[0] + st.local_shape[0] == (rank + 1) * per_shard + + +def _worker_save_then_load_offsets_symmetric(rank, world_size, port): + """Save-side and load-side ShardedTensors must produce identical offsets + and global_shape so DCP can correctly resharded between them. + + We don't run the real DCP save (avoids filesystem / async-writer issues + in CI); we just verify the symmetry property the load path relies on. + """ + update_gtp_config(pad_for_alignment=0) + dim0 = 16 + in_features = 4 + gtp_remat_group = _cached_new_group(list(range(world_size))) + + def _build(prefix): + weight = _make_gtp_shard(dim0, in_features, gtp_remat_group) + return make_sharded_tensors_for_checkpoint_with_gtp_remat( + {"weight": weight}, + prefix=prefix, + tensor_parallel_layers_axis_map={"weight": 0}, + sharded_offsets=(), + tp_group=_cached_new_group([rank]), + dp_cp_group=_cached_new_group(list(range(world_size))), + )["layer.weight"] + + save_st = _build("layer.") + load_st = _build("layer.") + assert save_st.global_shape == load_st.global_shape + assert save_st.global_offset == load_st.global_offset + assert save_st.local_shape == load_st.local_shape + assert save_st.replica_id == load_st.replica_id + + +def _worker_helper_offsets_ep_egtp(rank, world_size, port): + """EP=2, EGTP_remat=2 (4 ranks): routed-expert weight. + + Mirrors ``TEGroupedLinear.sharded_state_dict``: expert parallelism prepends a + global-expert axis through ``sharded_offsets``, and EGTP_remat shards each expert's + ``out_features`` (axis 0). The GTP_remat-aware checkpoint helper layers the EGTP_remat + axis-0 split on top of the prepended expert offset. + + rank → (ep_rank, egtp_rank): 0→(0,0) 1→(0,1) 2→(1,0) 3→(1,1). + """ + egtp_remat_group = _cached_new_group([0, 1]) if rank in (0, 1) else _cached_new_group([2, 3]) + + ep_size, egtp_remat_size, num_gemms = 2, 2, 1 + ep_rank = rank // 2 + egtp_rank = rank % 2 + per_expert_out = 4 + per_shard_out = per_expert_out // egtp_remat_size # 2 + in_features = 4 + num_global_experts = ep_size * num_gemms # 2 + global_expert_idx = ep_rank * num_gemms # + gemm_idx (0) + + weight = _make_gtp_shard(per_expert_out, in_features, egtp_remat_group) + assert weight.shape == ( + per_shard_out, + in_features, + ), f"rank={rank} EGTP_remat shape {tuple(weight.shape)} != ({per_shard_out}, {in_features})" + + sharded = make_sharded_tensors_for_checkpoint_with_gtp_remat( + {"weight": weight}, + prefix="", + tensor_parallel_layers_axis_map={"weight": 0}, + # EP prepends the global-expert axis; EGTP_remat shards out_features below it. + sharded_offsets=((0, global_expert_idx, num_global_experts),), + tp_group=_cached_new_group([rank]), # no TP in this case + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + st = sharded["weight"] + assert isinstance(st, ShardedTensor), f"Expected ShardedTensor, got {type(st)}" + # global shape = (num_global_experts, full_out_features, in_features) + assert st.global_shape == (num_global_experts, per_expert_out, in_features), ( + f"rank={rank} global_shape {st.global_shape} != " + f"({num_global_experts}, {per_expert_out}, {in_features})" + ) + # Prepended expert axis (axis 0): offset == this rank's global expert index. + assert ( + st.global_offset[0] == global_expert_idx + ), f"rank={rank} expert-axis offset {st.global_offset[0]} != {global_expert_idx}" + # EGTP_remat axis (weight axis 0, shifted to global axis 1): offset == egtp_rank · per_shard. + assert ( + st.global_offset[1] == egtp_rank * per_shard_out + ), f"rank={rank} EGTP_remat axis-1 offset {st.global_offset[1]} != {egtp_rank * per_shard_out}" + + +def _worker_helper_embedding_offsets(rank, world_size, port): + """Embedding / output_layer path: ``VocabParallelEmbedding.sharded_state_dict`` calls + ``make_tp_sharded_tensor_for_checkpoint`` DIRECTLY (it needs allow_shape_mismatch for + vocab padding), bypassing the GTP_remat-aware wrapper. So that helper itself must layer the + GTP_remat axis-0 split. TP=2, GTP_remat=2, tp_axis=0 → composite axis-0 offset, same as the + column-parallel case. + """ + gtp_remat_group = _cached_new_group([0, 1]) if rank in (0, 1) else _cached_new_group([2, 3]) + tp_group = _cached_new_group([0, 2]) if rank in (0, 2) else _cached_new_group([1, 3]) + + full_vocab, hidden = 8, 4 + tp_size, gtp_remat_size = 2, 2 + per_tp = full_vocab // tp_size # 4 + per_shard = per_tp // gtp_remat_size # 2 + + weight = _make_gtp_shard(per_tp, hidden, gtp_remat_group) + assert weight.shape == (per_shard, hidden) + + st = make_tp_sharded_tensor_for_checkpoint( + tensor=weight, + key="embedding.word_embeddings.weight", + tp_axis=0, + allow_shape_mismatch=True, # how VocabParallelEmbedding calls it + prepend_offsets=(), + tp_group=tp_group, + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + assert isinstance(st, ShardedTensor), f"Expected ShardedTensor, got {type(st)}" + tp_rank = rank // 2 + gtp_rank = rank % 2 + expected_offset = (tp_rank * gtp_remat_size + gtp_rank) * per_shard + assert ( + st.global_offset[0] == expected_offset + ), f"rank={rank} embedding axis-0 offset {st.global_offset[0]} != {expected_offset}" + assert ( + st.global_shape[0] == full_vocab + ), f"rank={rank} embedding global axis-0 {st.global_shape[0]} != {full_vocab}" + + +def _worker_helper_public_wrapper_delegates(rank, world_size, port): + """The public ``make_sharded_tensors_for_checkpoint`` (the entry point most layers call, + e.g. ColumnParallelLinear / output_layer) must detect a GTPShardedParam and produce the + GTP_remat-composite offset — i.e. it delegates to the GTP_remat-aware path not the vanilla + TP-only one. TP=2, GTP_remat=2, column-parallel (tp_axis=0). + """ + gtp_remat_group = _cached_new_group([0, 1]) if rank in (0, 1) else _cached_new_group([2, 3]) + tp_group = _cached_new_group([0, 2]) if rank in (0, 2) else _cached_new_group([1, 3]) + + full_out, in_features = 8, 4 + tp_size, gtp_remat_size = 2, 2 + per_tp_out = full_out // tp_size # 4 + per_shard_out = per_tp_out // gtp_remat_size # 2 + + weight = _make_gtp_shard(per_tp_out, in_features, gtp_remat_group) + + sharded = make_sharded_tensors_for_checkpoint( + {"weight": weight}, + prefix="layer.", + tensor_parallel_layers_axis_map={"weight": 0}, + sharded_offsets=(), + tp_group=tp_group, + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + st = sharded["layer.weight"] + assert isinstance(st, ShardedTensor), f"Expected ShardedTensor, got {type(st)}" + tp_rank = rank // 2 + gtp_rank = rank % 2 + expected_offset = (tp_rank * gtp_remat_size + gtp_rank) * per_shard_out + assert st.global_offset[0] == expected_offset, ( + f"rank={rank} public wrapper did not produce the GTP_remat-composite offset: " + f"{st.global_offset[0]} != {expected_offset} (delegation to the GTP_remat path failed?)" + ) + assert ( + st.global_shape[0] == full_out + ), f"rank={rank} global axis-0 {st.global_shape[0]} != {full_out}" + + +def _worker_helper_replicated_sink_rejects_gtp(rank, world_size, port): + """Sanity guard: a GTPShardedParam must NEVER be saved via the replicated + make_sharded_tensor_for_checkpoint (it would record a shard-sized global shape). + The helper asserts; this pins that behaviour. + """ + from megatron.core.utils import make_sharded_tensor_for_checkpoint + + gtp_remat_group = _cached_new_group([0, 1]) if rank in (0, 1) else _cached_new_group([2, 3]) + weight = _make_gtp_shard(4, 4, gtp_remat_group) + with pytest.raises(AssertionError): + make_sharded_tensor_for_checkpoint( + weight, + "weight", + tp_group=_cached_new_group([rank]), + dp_cp_group=_cached_new_group(list(range(world_size))), + ) + + +def _worker_mamba_replicated_param_replica_ids(rank, world_size, port): + """MambaMixer.sharded_state_dict under GTP_remat: replicated directly-owned params + (A_log / dt_bias / D / conv1d.*) must get conflict-free replica_ids -- unique across the + peers holding each chunk, exactly one writer -- so DCP elects a single writer per chunk. + """ + GTP_remat = 2 # world=4 -> tp1 * gtp2 * dp2 (exercises both gtp_remat peers and replicate DP) + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=GTP_remat + ) + model_parallel_cuda_manual_seed(42) + pg = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp', 'gtp_remat']) + + config = TransformerConfig( + num_attention_heads=32, + num_layers=1, + hidden_size=4096, + mamba_num_heads=128, + mamba_head_dim=64, + mamba_state_dim=128, + mamba_num_groups=8, + use_mamba_mem_eff_path=True, + params_dtype=torch.bfloat16, + hidden_dropout=0.0, + bias_dropout_fusion=False, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + ) + submodules = MambaLayerSubmodules( + mixer=ModuleSpec( + module=MambaMixer, + submodules=MambaMixerSubmodules( + in_proj=TELayerNormColumnParallelLinear, out_proj=TERowParallelLinear + ), + ), + mamba_bda=get_bias_dropout_add, + ) + layer = MambaLayer(config, submodules, layer_number=1, pg_collection=pg).cuda() + assert any( + isinstance(p, GTPShardedParam) for p in layer.parameters() + ), "GTP_remat not active: no GTPShardedParam in the GTP_remat=2 Mamba layer" + + # Checkpoint replica election for gtp_remat-REPLICATED params needs the gtp_remat-INCLUSIVE + # group so gtp_remat peers get distinct replica_ids (matches production's get_default + # metadata, which uses the gtp_remat-inclusive default). The replicate group would collide. + metadata = {'dp_cp_group': ps.get_data_parallel_group(with_context_parallel=True)} + sd = layer.mixer.sharded_state_dict(prefix='mixer.', metadata=metadata) + + target_bases = {'A_log', 'dt_bias', 'D', 'conv1d.weight', 'conv1d.bias'} + local = {} + for key, val in sd.items(): + base = key.split('mixer.', 1)[-1] + if base in target_bases and isinstance( + val, (ShardedTensor, ShardedTensorFactory, ShardedObject) + ): + rid = val.replica_id + if isinstance(rid, tuple): + local[base] = tuple(rid) + + gathered = [None] * world_size + dist.all_gather_object(gathered, local) + + ps.destroy_model_parallel() + ps.initialize_model_parallel() + GTPShardedParam._chain_state = {} + + if rank == 0: + bases = set(gathered[0]) + assert bases, "no GTP_remat-replicated tiny params found in MambaMixer sharded_state_dict" + for base in sorted(bases): + rids = [g[base] for g in gathered] + assert ( + len(set(rids)) == world_size + ), f"{base}: replica_id collision across ranks -> DCP write conflict: {rids}" + n_writers = sum(is_main_replica(r) for r in rids) + assert n_writers == 1, f"{base}: expected exactly 1 writer, got {n_writers}: {rids}" + + +def _worker_replicated_param_needs_gtp_inclusive_dp_cp(rank, world_size, port): + """Regression for the checkpoint-save duplicate-writer bug in save_checkpoint_and_time. + + A REPLICATED param's replica_id must use the gtp_remat-INCLUSIVE group (``pg.dp_cp_gtp_remat``); + the gtp-excluded ``pg.dp_cp`` collapses gtp_remat peers to one replica_id -> multiple writers + -> save validation failure. world=4 -> tp1*gtp2*dp2 (replicate=2 ranks, inclusive=4). + """ + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + pg = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=['tp', 'dp_cp', 'dp_cp_gtp_remat'] + ) + + # The two attributes save_checkpoint_and_time may read must differ under GTP_remat, else the + # group choice would be moot: full = replicate x gtp_remat(2). + assert ( + get_pg_size(pg.dp_cp_gtp_remat) == get_pg_size(pg.dp_cp) * 2 + ), f"full={get_pg_size(pg.dp_cp_gtp_remat)} replicate={get_pg_size(pg.dp_cp)}" + + replicated = torch.nn.Parameter(torch.zeros(8, 4, dtype=torch.bfloat16, device="cuda")) + + def _gather_replica_ids(dp_cp_group): + sd = make_sharded_tensors_for_checkpoint( + {"w": replicated}, + prefix="", + tensor_parallel_layers_axis_map={}, + tp_group=pg.tp, + dp_cp_group=dp_cp_group, + ) + out = [None] * world_size + dist.all_gather_object(out, tuple(sd["w"].replica_id)) + return out + + rids_replicate = _gather_replica_ids(pg.dp_cp) # the bug + rids_full = _gather_replica_ids(pg.dp_cp_gtp_remat) # the fix + + ps.destroy_model_parallel() + ps.initialize_model_parallel() + + if rank == 0: + # Replicate (gtp-excluded) group: gtp_remat peers collapse -> >1 writer (reproduces bug). + assert ( + sum(is_main_replica(r) for r in rids_replicate) > 1 + ), f"replicate dp_cp should collide across gtp_remat peers, got {rids_replicate}" + # gtp_remat-inclusive group: every holder distinct, exactly one writer. + assert len(set(rids_full)) == world_size, f"full-group replica_id collision: {rids_full}" + assert ( + sum(is_main_replica(r) for r in rids_full) == 1 + ), f"full group must elect exactly one writer, got {rids_full}" + + +def _worker_embedding_writer_election_gtp_inclusive_default(rank, world_size, port): + """VocabParallelEmbedding calls make_tp_sharded_tensor_for_checkpoint directly (needs + allow_shape_mismatch), so its GTP writer election must use the gtp_remat-EXCLUDED DP group. + Assert every axis-0 offset has exactly one main-replica writer and the offsets tile the vocab. + + world=4 -> tp1 * gtp2 * dp2: vocab split in 2 (gtp), each half replicated on 2 dp ranks. + """ + from collections import defaultdict + + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + gtp_group = ps.get_gtp_weight_remat_group() + assert gtp_group.size() == 2, f"expected gtp_remat_size=2, got {gtp_group.size()}" + + full_vocab, hidden = 8, 4 + per_shard = full_vocab // gtp_group.size() # 4 (tp=1 -> per_tp == full_vocab) + weight = _make_gtp_shard(full_vocab, hidden, gtp_group) + assert weight.shape == (per_shard, hidden) + + st = make_tp_sharded_tensor_for_checkpoint( + tensor=weight, + key="embedding.word_embeddings.weight", + tp_axis=0, + allow_shape_mismatch=True, # how VocabParallelEmbedding calls it + prepend_offsets=(), + tp_group=ps.get_tensor_model_parallel_group(), + dp_cp_group=ps.get_data_parallel_group(with_context_parallel=True), + ) + mine = (int(st.global_offset[0]), tuple(st.replica_id)) + gathered = [None] * world_size + dist.all_gather_object(gathered, mine) + + ps.destroy_model_parallel() + ps.initialize_model_parallel() + GTPShardedParam._chain_state = {} + + if rank == 0: + by_offset = defaultdict(list) + for off, rid in gathered: + by_offset[off].append(rid) + assert set(by_offset) == {0, per_shard}, f"vocab offsets must tile: {sorted(by_offset)}" + for off, rids in by_offset.items(): + n_writers = sum(is_main_replica(r) for r in rids) + assert n_writers == 1, ( + f"vocab offset {off}: expected exactly 1 checkpoint writer, got {n_writers} " + f"(replica_ids {rids}); a gtp_remat-inclusive default DP group leaves a shard " + f"with no main-replica writer -> 'Invalid access pattern' at save" + ) + + +def _worker_mamba_inproj_optim_param_map(rank, world_size, port): + """GTP_remat+Muon ckpt fix: in_proj's gathered+split model entry does NOT id-match the + per-shard optimizer param, so get_param_id_to_sharded_param_map misses it (the KeyError seen in + Float16OptimizerWithFloat16Params.sharded_state_dict). Verify the per-shard fallback used by the + fix restores a ShardedTensor with local_shape == the optimizer param shape, which + make_sharded_optimizer_tensor then accepts. + """ + from megatron.core.dist_checkpointing.optimizer import ( + get_param_id_to_sharded_param_map, + make_sharded_optimizer_tensor, + ) + from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( + tag_gtp_params_with_names, + ) + + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + model_parallel_cuda_manual_seed(42) + pg = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp', 'gtp_remat']) + config = TransformerConfig( + num_attention_heads=32, + num_layers=1, + hidden_size=4096, + mamba_num_heads=128, + mamba_head_dim=64, + mamba_state_dim=128, + mamba_num_groups=8, + use_mamba_mem_eff_path=True, + params_dtype=torch.bfloat16, + hidden_dropout=0.0, + bias_dropout_fusion=False, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + ) + submodules = MambaLayerSubmodules( + mixer=ModuleSpec( + module=MambaMixer, + submodules=MambaMixerSubmodules( + in_proj=TELayerNormColumnParallelLinear, out_proj=TERowParallelLinear + ), + ), + mamba_bda=get_bias_dropout_add, + ) + layer = MambaLayer(config, submodules, layer_number=1, pg_collection=pg).cuda() + tag_gtp_params_with_names(layer) # set _debug_name (mirrors production setup) + + in_proj_w = layer.mixer.in_proj.weight + assert isinstance(in_proj_w, GTPShardedParam), "in_proj.weight should be GTP_remat-sharded" + + metadata = {'dp_cp_group': ps.get_data_parallel_group(with_context_parallel=True)} + model_sd = layer.mixer.sharded_state_dict(prefix='mixer.', metadata=metadata) + + # Reproduce the gap: in_proj's per-shard optim param has no id-match in the model dict. + id_map = get_param_id_to_sharded_param_map(model_sd, [in_proj_w]) + assert 0 not in id_map, "expected in_proj to be MISSING from id map (the KeyError gap)" + + # The fix's per-shard fallback restores a matching entry. + key = in_proj_w._debug_name or '_gtp_optim_param_0' + entry = make_sharded_tensors_for_checkpoint_with_gtp_remat( + {key: in_proj_w}, + prefix='', + tensor_parallel_layers_axis_map={key: 0}, + tp_group=ps.get_tensor_model_parallel_group(), + dp_cp_group=ps.get_data_parallel_group(with_context_parallel=True), + )[key] + assert tuple(entry.local_shape) == tuple(in_proj_w.shape), ( + f"per-shard entry local_shape {tuple(entry.local_shape)} != param shape " + f"{tuple(in_proj_w.shape)}" + ) + # make_sharded_optimizer_tensor must accept it for a same-shape optimizer state tensor. + opt_state = torch.zeros_like(in_proj_w) + osh = make_sharded_optimizer_tensor(entry, opt_state, prefix='optimizer.state.exp_avg') + assert osh is not None + + ps.destroy_model_parallel() + ps.initialize_model_parallel() + GTPShardedParam._chain_state = {} + + +def _worker_save_load_roundtrip_needs_gtp_inclusive_group(rank, world_size, ckpt_base): + """Save->load roundtrip: save and load must use the gtp_remat-INCLUSIVE replica group. + + Loading with the gtp_remat-EXCLUDING ``pg.dp_cp`` collides replica_ids across gtp_remat + peers -> DCP 'Invalid access pattern' (the a55b load failure). world=4 -> tp1*gtp2*dp2. + """ + from megatron.core.dist_checkpointing import load, save + from tests.unit_tests.dist_checkpointing import TempNamedDir + + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + try: + pg = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=['tp', 'dp_cp', 'dp_cp_gtp_remat'] + ) + # The two group choices must differ under GTP_remat, else the test is moot. + assert ( + get_pg_size(pg.dp_cp_gtp_remat) == get_pg_size(pg.dp_cp) * 2 + ), f"full={get_pg_size(pg.dp_cp_gtp_remat)} replicate={get_pg_size(pg.dp_cp)}" + + # A GTP-replicated param: byte-identical on every rank (like decoder.final_norm.weight). + replicated = torch.nn.Parameter( + torch.arange(32, dtype=torch.bfloat16, device="cuda").reshape(8, 4) + ) + + def _sd(dp_cp_group): + return make_sharded_tensors_for_checkpoint( + {"w": replicated}, + prefix="", + tensor_parallel_layers_axis_map={}, + tp_group=pg.tp, + dp_cp_group=dp_cp_group, + ) + + # Negative invariant (uniform, no failed collective): under the gtp_remat-EXCLUDING + # replicate group the gtp_remat peers collapse to the same replica_id -> >1 main writer + # for the identical element. That is exactly what makes DCP load-validation raise the + # a55b 'Invalid access pattern'. (validate_sharding_integrity raises on rank 0 only, so a + # real failed load() can't be asserted cleanly across ranks; we assert the root condition.) + rids_excl = [None] * world_size + dist.all_gather_object(rids_excl, tuple(_sd(pg.dp_cp)["w"].replica_id)) + rids_incl = [None] * world_size + dist.all_gather_object(rids_incl, tuple(_sd(pg.dp_cp_gtp_remat)["w"].replica_id)) + if rank == 0: + assert ( + sum(is_main_replica(r) for r in rids_excl) > 1 + ), f"gtp-excluding dp_cp must collide across gtp_remat peers (the bug): {rids_excl}" + assert ( + sum(is_main_replica(r) for r in rids_incl) == 1 + ), f"gtp-inclusive group must elect exactly one writer: {rids_incl}" + + # Positive end-to-end roundtrip through the real DCP save/load with the gtp_remat-inclusive + # group (what save_checkpoint_and_time and the fixed load_checkpoint both thread): save and + # load must agree on this group, and the replicated data must round-trip intact. + with TempNamedDir(ckpt_base / 'gtp_dcp_roundtrip', sync=True) as ckpt_dir: + save(_sd(pg.dp_cp_gtp_remat), ckpt_dir) + loaded = load(_sd(pg.dp_cp_gtp_remat), ckpt_dir) + assert torch.equal(loaded["w"].cpu(), replicated.detach().cpu()), loaded["w"] + finally: + ps.destroy_model_parallel() + ps.initialize_model_parallel() + + +# --------------------------------------------------------------------------- +# Test class wrappers (4-GPU) +# --------------------------------------------------------------------------- + + +@pytest.mark.run_only_on_devices_with_compute_capability(compute_capability=(10, 0)) +class TestGtpDcpHelper: + def test_mamba_replicated_param_replica_ids(self): + _require_world_size(4) + _worker_mamba_replicated_param_replica_ids(dist.get_rank(), 4, None) + + def test_mamba_inproj_optim_param_map(self): + _require_world_size(4) + _worker_mamba_inproj_optim_param_map(dist.get_rank(), 4, None) + + def test_replicated_param_needs_gtp_inclusive_dp_cp(self): + _require_world_size(4) + _worker_replicated_param_needs_gtp_inclusive_dp_cp(dist.get_rank(), 4, None) + + def test_save_load_roundtrip_needs_gtp_inclusive_group(self, tmp_path_dist_ckpt): + _require_world_size(4) + _worker_save_load_roundtrip_needs_gtp_inclusive_group( + dist.get_rank(), 4, tmp_path_dist_ckpt + ) + + def test_composite_offset_same_axis(self): + _require_world_size(4) + _worker_helper_offsets_tp_eq_gtp_axis(dist.get_rank(), 4, None) + + def test_native_fp8_dcp_save(self): + _require_world_size(4) + _worker_native_fp8_dcp_save(dist.get_rank(), 4, None) + + def test_native_fp8_dcp_load_copy(self): + _require_world_size(4) + _worker_native_fp8_dcp_load_copy(dist.get_rank(), 4, None) + + def test_dual_offsets_cross_axis(self): + _require_world_size(4) + _worker_helper_offsets_tp_neq_gtp_axis(dist.get_rank(), 4, None) + + def test_ep_egtp_offsets(self): + _require_world_size(4) + _worker_helper_offsets_ep_egtp(dist.get_rank(), 4, None) + + def test_embedding_offsets(self): + _require_world_size(4) + _worker_helper_embedding_offsets(dist.get_rank(), 4, None) + + def test_embedding_writer_election(self): + _require_world_size(4) + _worker_embedding_writer_election_gtp_inclusive_default(dist.get_rank(), 4, None) + + def test_public_wrapper_delegates(self): + _require_world_size(4) + _worker_helper_public_wrapper_delegates(dist.get_rank(), 4, None) + + def test_replicated_sink_rejects_gtp(self): + _require_world_size(4) + _worker_helper_replicated_sink_rejects_gtp(dist.get_rank(), 4, None) + + def test_no_op_no_gtp_remat(self): + _require_world_size(4) + _worker_helper_no_op_no_gtp_remat(dist.get_rank(), 4, None) + + def test_inproj_no_pad(self): + _require_world_size(4) + _worker_helper_padded_inproj_no_pad_case(dist.get_rank(), 4, None) + + def test_inproj_with_pad(self): + _require_world_size(4) + _worker_helper_padded_inproj_pad_case(dist.get_rank(), 4, None) + + def test_cross_topology_reshard_metadata(self): + _require_world_size(4) + _worker_helper_cross_topology_reshard_metadata(dist.get_rank(), 4, None) + + def test_save_then_load_offsets_symmetric(self): + _require_world_size(4) + _worker_save_then_load_offsets_symmetric(dist.get_rank(), 4, None) diff --git a/tests/unit_tests/generalized_tensor_parallel/test_gtp_fp8_param_gather.py b/tests/unit_tests/generalized_tensor_parallel/test_gtp_fp8_param_gather.py new file mode 100644 index 00000000000..127a2fc1761 --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_gtp_fp8_param_gather.py @@ -0,0 +1,201 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""GTP + MXFP8 --fp8-param-gather / --reuse-grad-buf-for-mxfp8-param-ag correctness. + +Asserts the two MXFP8 param-gather knobs don't change training: a GTP (weight-remat=2) loss +trajectory with the knobs on must match the same run with them off. Reuses the full DDP + +DistributedOptimizer harness from ``test_fp8_param.py::TestFP8Param`` by composition (imported +under a non-``Test*`` alias so pytest doesn't re-collect it), flipping GTP on via +``tensor_parallel_num_weight_shards`` (= tp x gtp_weight_remat_size). +""" + +import pytest +import torch + +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + +if not HAVE_GTP: + pytest.skip("GTP requires TransformerEngine >= 2.19", allow_module_level=True) + +from megatron.core.utils import is_te_min_version +from megatron.training.utils import get_device_arch_version + +# Non-"Test*" alias so pytest does not re-collect the whole TestFP8Param suite here (wrong +# world/DP config + global-state pollution); reused by composition only. +from tests.unit_tests.test_fp8_param import TestFP8Param as _FP8ParamHarness +from tests.unit_tests.test_fp8_param import fp8_available, reason_for_no_fp8 + + +class TestGTPFp8ParamGather: + """GTP weight-remat=2 loss-trajectory parity for the MXFP8 param-gather knobs.""" + + @pytest.mark.skipif( + get_device_arch_version() < 10, reason="MXFP8 is supported since Blackwell architecture" + ) + @pytest.mark.skipif(not fp8_available, reason=reason_for_no_fp8) + @pytest.mark.skipif(not is_te_min_version("2.3.0.dev0"), reason="TE 2.3.0.dev0 is required") + @pytest.mark.parametrize("dp_overlap", [(False, False), (True, True)]) + # (tp_size, num_weight_shards, min_gpus): tp2 case guards a TP/GTP axis-order inversion in + # native-FP8 init (real TE TP divide + sequence_parallel x GTP2). + @pytest.mark.parametrize("tp_case", [(1, 2, 2), (2, 4, 4)]) + def test_gtp_mxfp8_fp8_param_gather(self, dp_overlap, tp_case): + """GTP weight-remat=2: fp8-param loss must track pure-BF16 loss within MXFP8 noise. + + A frozen fp8 forward weight (optimizer updates not reaching the native fp8 shard) instead + leaves fp8 flat while BF16 descends (~1+ gap). dp_overlap=(overlap_param_gather, + overlap_grad_reduce); the overlap leg exercises the ``_copy_main_params_to_param_buffer`` + path GTP hooks for --reuse-grad-buf-for-mxfp8-param-ag. + """ + tp_size, num_shards, min_gpus = tp_case + if torch.cuda.device_count() < min_gpus: + pytest.skip(f"Requires {min_gpus} CUDA devices for TP{tp_size} x GTP weight-remat=2") + + harness = _FP8ParamHarness() + harness.setup_method(None) + # num-microbatches uses data_parallel_size = world/tp (gtp is a DP sub-axis). + harness.micro_batch_size = 1 + try: + common = dict( + tp_size=tp_size, + global_batch_size=4, + overlap_param_gather=dp_overlap[0], + overlap_grad_reduce=dp_overlap[1], + tensor_parallel_num_weight_shards=num_shards, # tp * N => gtp_weight_remat_size=N + # Untie: the tied path feeds the GTP-sharded embedding into a Megatron-native + # ColumnParallelLinear, which does no GTP all-gather (TE-only) and fails its check. + untie_embeddings_and_output_weights=True, + ) + loss_fp8 = harness._run_test_helper(recipe="mxfp8", fp8_param_gather=True, **common) + # Pure BF16 GTP reference: fp8=None overrides the harness default (recipe inert). + loss_bf16 = harness._run_test_helper( + recipe="delayed", fp8_param_gather=False, fp8=None, **common + ) + # Max drift ~0.03 over 100 steps (MXFP8 noise); 0.05 stays above it and trips on the + # ~1+ frozen-weight gap. + diff = (loss_fp8 - loss_bf16).abs().max().item() + assert diff < 0.05, ( + f"GTP+mxfp8 fp8-param-gather loss diverges from pure-BF16 GTP baseline " + f"(max per-step |diff|={diff:.4f}; fp8: {loss_fp8[0]:.3f}->{loss_fp8[-1]:.3f}, " + f"bf16: {loss_bf16[0]:.3f}->{loss_bf16[-1]:.3f})." + ) + finally: + harness.teardown_method(None) + # Restore GTP_CONFIG defaults mutated by the mxfp8 arg setup. + from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( + update_gtp_config, + ) + + update_gtp_config(pad_for_alignment=16, calculate_per_token_loss=False) + + @pytest.mark.skipif( + get_device_arch_version() < 10, reason="MXFP8 is supported since Blackwell architecture" + ) + @pytest.mark.skipif(not fp8_available, reason=reason_for_no_fp8) + @pytest.mark.skipif(not is_te_min_version("2.3.0.dev0"), reason="TE 2.3.0.dev0 is required") + def test_gtp_mxfp8_moe_fp8_param_gather(self): + """MoE grouped-expert (TEGroupedLinear) native-FP8 GTP: loss must track pure-BF16 GTP. + + Covers the EGTP-sharded expert weights built as native MXFP8 shards under + --fp8-param-gather — the gap the dense test (attention + dense MLP) leaves. Same parity + assertion as the dense case. + """ + if torch.cuda.device_count() < 4: + pytest.skip("Requires 4 CUDA devices for EP=2 x GTP weight-remat=2 MoE") + + harness = _FP8ParamHarness() + harness.setup_method(None) + harness.micro_batch_size = 1 + try: + common = dict( + tp_size=1, + global_batch_size=4, + overlap_param_gather=True, + overlap_grad_reduce=True, + tensor_parallel_num_weight_shards=2, # tp=1 * 2 => gtp_weight_remat_size=2 (EGTP) + untie_embeddings_and_output_weights=True, + # MoE grouped experts (mirror test_mxfp8_moe), EP=2. + num_experts=2, + moe_grouped_gemm=True, + expert_model_parallel_size=2, + moe_token_dispatcher_type="alltoall", + moe_router_topk=1, + moe_router_pre_softmax=True, + moe_router_load_balancing_type="none", + moe_aux_loss_coeff=0.0, + moe_ffn_hidden_size=128, + ) + loss_fp8 = harness._run_test_helper(recipe="mxfp8", fp8_param_gather=True, **common) + loss_bf16 = harness._run_test_helper( + recipe="delayed", fp8_param_gather=False, fp8=None, **common + ) + diff = (loss_fp8 - loss_bf16).abs().max().item() + assert diff < 0.05, ( + f"GTP+mxfp8 MoE fp8-param-gather loss diverges from pure-BF16 GTP baseline " + f"(max per-step |diff|={diff:.4f}; fp8: {loss_fp8[0]:.3f}->{loss_fp8[-1]:.3f}, " + f"bf16: {loss_bf16[0]:.3f}->{loss_bf16[-1]:.3f})." + ) + finally: + harness.teardown_method(None) + from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( + update_gtp_config, + ) + + update_gtp_config(pad_for_alignment=16, calculate_per_token_loss=False) + + @pytest.mark.skipif( + get_device_arch_version() < 10, reason="MXFP8 is supported since Blackwell architecture" + ) + @pytest.mark.skipif(not fp8_available, reason=reason_for_no_fp8) + @pytest.mark.skipif(not is_te_min_version("2.3.0.dev0"), reason="TE 2.3.0.dev0 is required") + def test_gtp_mxfp8_save_does_not_perturb_training(self): + """A checkpoint save must NOT mutate the live weights. + + Runs GTP+mxfp8+fp8-param-gather twice with identical seeds — once driving the production + save path mid-training (force_param_sync + sharded_state_dict), once without — and requires + matching loss trajectories. overlap_param_gather=True makes should_disable_forward_pre_hook + True so force_param_sync actually runs; passing the optimizer copies FP32 masters into the + param buffer first, so the copy-back re-quantizes the GTP native-FP8 shard from masters (not + stale grad scratch). Guards the historical post-save loss spike (seen at a55b) — a save + side-effect test_gtp_dcp can't see (it never trains after saving). + """ + if torch.cuda.device_count() < 2: + pytest.skip("Requires at least 2 CUDA devices for GTP weight-remat=2") + + common = dict( + tp_size=1, + recipe="mxfp8", + fp8_param_gather=True, + overlap_param_gather=True, + overlap_grad_reduce=True, + global_batch_size=4, + tensor_parallel_num_weight_shards=2, + untie_embeddings_and_output_weights=True, + ) + try: + h1 = _FP8ParamHarness() + h1.setup_method(None) + h1.micro_batch_size = 1 + loss_baseline = h1._run_test_helper(**common) + h1.teardown_method(None) + + h2 = _FP8ParamHarness() + h2.setup_method(None) + h2.micro_batch_size = 1 + loss_saved = h2._run_test_helper(save_at_steps=(5, 10, 15), **common) + h2.teardown_method(None) + + diff = (loss_baseline - loss_saved).abs() + worst = diff.max().item() + # Save runs a real MXFP8 force_param_sync (not bit-exact vs no-save), so allow re-gather + # noise (~0.03/100 steps); 0.1 clears it and still catches the pre-fix O(10) spike. + assert worst < 0.1, ( + f"Checkpoint save perturbed training (max per-step |diff|={worst:.4f} at step " + f"{int(diff.argmax())}); the forced pre-save param-sync is corrupting live FP8 " + f"weights. saved-run around first save: {loss_saved[4:8].tolist()}" + ) + finally: + from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( + update_gtp_config, + ) + + update_gtp_config(pad_for_alignment=16, calculate_per_token_loss=False) diff --git a/tests/unit_tests/generalized_tensor_parallel/test_gtp_grad_correctness.py b/tests/unit_tests/generalized_tensor_parallel/test_gtp_grad_correctness.py new file mode 100644 index 00000000000..b969e5e5bf8 --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_gtp_grad_correctness.py @@ -0,0 +1,559 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Numeric repro: GTP_remat gradient correctness through the REAL +DDP + distributed-optimizer + finalize path, with replicate (DP) > 1. + +The validated loss-trajectory test uses DP=1 (replicate=1) and manual +SGD on main_grad, so it cannot catch a gradient-reduction error that only shows +up when the dist-opt shards over a replicate group of size > 1 (the new-at-64-GPU +condition: DP2 x GTP16). This test reproduces that condition at small scale +(world=4 = GTP2 x DP2) and checks the gradient end-to-end against a trusted +no-GTP_remat DP=4 baseline. + +Decisive choices: + * SGD lr=1.0 (NOT Adam): the step is scale-SENSITIVE, so a gtp_remat x gradient + under-scale shows up directly as a gtp_remat x smaller weight delta. Adam would + normalize a uniform scale error away and mask the bug. + * Distinct input per rank (seed=rank): each data-parallel position sees a + different batch (the HSDP guarantee), so the correct reduced grad is the + MEAN over all 4 positions. Baseline (DP4) and GTP_remat (GTP2xDP2) both + span the same 4 positions, so their reduced grads -- and thus post-step + weights and grad-norm -- must match. +""" + +import pytest +import torch +import torch.distributed as dist + +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + +if not HAVE_GTP: + pytest.skip("GTP requires TransformerEngine >= 2.19", allow_module_level=True) + +from megatron.core.tensor_parallel.generalized_tensor_parallelism import GTPShardedParam +from tests.unit_tests.generalized_tensor_parallel.gtp_test_utils import ( # noqa: F401 + _run_distributed, + _torchrun_dist_init, + reset_fp8_state, + reset_gtp_globals, +) + +HIDDEN = 256 +NUM_HEADS = 8 +FFN_HIDDEN = 512 +NUM_LAYERS = 1 +SEQ = 16 +BATCH = 1 +LR = 1.0 # scale-sensitive SGD step +dtype = torch.bfloat16 + + +def _make_config(calculate_per_token_loss=False): + from megatron.core.transformer.transformer_config import TransformerConfig + + return TransformerConfig( + num_attention_heads=NUM_HEADS, + num_layers=NUM_LAYERS, + hidden_size=HIDDEN, + ffn_hidden_size=FFN_HIDDEN, + add_bias_linear=False, + params_dtype=dtype, + hidden_dropout=0.0, + attention_dropout=0.0, + bias_dropout_fusion=False, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + calculate_per_token_loss=calculate_per_token_loss, + ) + + +def _make_stack(config, pg_collection): + from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec + + spec = get_gpt_layer_with_transformer_engine_spec() + return torch.nn.ModuleList( + [ + spec.module(config, spec.submodules, layer_number=i + 1, pg_collection=pg_collection) + for i in range(NUM_LAYERS) + ] + ) + + +def _build_ddp(stack, calculate_per_token_loss=False): + """Wrap the stack in a NON-distributed-optimizer DDP so main_grad holds the + full all-reduced gradient (no optimizer needed; no Adam scale-invariance to + mask a scaling error).""" + from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig + + config = _make_config(calculate_per_token_loss=calculate_per_token_loss) + ddp_config = DistributedDataParallelConfig( + use_distributed_optimizer=False, overlap_grad_reduce=False + ) + module = torch.nn.Sequential() + for i, layer in enumerate(stack): + module.add_module(str(i), layer) + return DistributedDataParallel(config, ddp_config, module) + + +def _run_one_backward(ddp_model, rank, calculate_per_token_loss=False): + ddp_model.zero_grad_buffer() + # Distinct input per rank => the correct reduced grad is the MEAN over ranks. + torch.manual_seed(1000 + rank) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + out = x + for layer in ddp_model.module.children(): + out, _ = layer(out, attention_mask=None) + loss = out.float().mean() + loss.backward() + # Sync ONCE: finish_grad_sync() triggers the (single) grad reduction for + # overlap_grad_reduce=False. Do NOT also call start_grad_sync() — that double- + # reduces, which is idempotent at full-DP size but halves at replicate size. + ddp_model.finish_grad_sync() + from megatron.core.distributed.finalize_model_grads import ( + _allreduce_replicated_grads_over_gtp_remat_group, + ) + + _allreduce_replicated_grads_over_gtp_remat_group( + [ddp_model], calculate_per_token_loss=calculate_per_token_loss + ) + return float(loss.item()) + + +def _full_main_grads(stack): + """Reconstruct full (unsharded) reduced gradients keyed by param name. + + GTPShardedParam.main_grad is the local gtp_remat shard -> all-gather over the gtp_remat + group. Non-GTP_remat params are replicated -> take the local (already gtp_remat-summed) copy. + """ + from megatron.core import parallel_state as ps + + out = {} + for layer in stack: + for name, p in layer.named_parameters(): + g_attr = 'main_grad' if hasattr(p, 'main_grad') else 'grad' + mg = getattr(p, g_attr) + if isinstance(p, GTPShardedParam): + g = ps.get_gtp_weight_remat_group() + shards = [torch.empty_like(mg) for _ in range(g.size())] + dist.all_gather(shards, mg.contiguous(), group=g) + out[name] = torch.cat(shards, dim=0).float().cpu() + else: + out[name] = mg.detach().float().cpu() + return out + + +def _worker(rank, world_size, port, calculate_per_token_loss=False): + from megatron.core import parallel_state as ps + from megatron.core.process_groups_config import ProcessGroupCollection + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + + # ---------- Phase A: baseline, GTP_remat=1 DP=4 (trusted standard path) ---------- + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=1 + ) + model_parallel_cuda_manual_seed(42) + pgc = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp', 'gtp_remat']) + base_stack = _make_stack(_make_config(calculate_per_token_loss=calculate_per_token_loss), pgc) + for layer in base_stack: + layer.cuda() + for p in base_stack.parameters(): + dist.broadcast(p.data, src=0) + saved = {n: p.data.clone() for n, p in base_stack.named_parameters()} + + base_ddp = _build_ddp(base_stack, calculate_per_token_loss=calculate_per_token_loss) + _run_one_backward(base_ddp, rank, calculate_per_token_loss=calculate_per_token_loss) + base_grads = _full_main_grads(base_stack) + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + + # ---------- Phase B: GTP_remat=2 DP=2 (replicate>1!) ---------- + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + model_parallel_cuda_manual_seed(42) + pgc = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp', 'gtp_remat']) + gtp_stack = _make_stack(_make_config(calculate_per_token_loss=calculate_per_token_loss), pgc) + for layer in gtp_stack: + layer.cuda() + + g = ps.get_gtp_weight_remat_group() + gtp_rank = g.rank() + assert g.size() == 2, f"expected gtp_remat shard group size 2, got {g.size()}" + + # Load the SAME init weights as baseline: GTP_remat params get their gtp_remat shard. + for name, p in gtp_stack.named_parameters(): + full = saved[name] + if isinstance(p, GTPShardedParam): + ss = p.shape[0] + p.data.copy_(full[gtp_rank * ss : (gtp_rank + 1) * ss]) + else: + p.data.copy_(full) + + gtp_ddp = _build_ddp(gtp_stack, calculate_per_token_loss=calculate_per_token_loss) + _run_one_backward(gtp_ddp, rank, calculate_per_token_loss=calculate_per_token_loss) + gtp_grads = _full_main_grads(gtp_stack) + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + + # ---------- Compare reduced gradients on rank 0 ---------- + if rank == 0: + max_err = 0.0 + worst = None + for name in base_grads: + bg, gg = base_grads[name], gtp_grads[name] + assert bg.shape == gg.shape, f"{name}: {bg.shape} vs {gg.shape}" + err = (bg - gg).abs().max().item() + denom = bg.abs().max().item() + 1e-8 + rel = err / denom + ratio = (gg.norm() / (bg.norm() + 1e-12)).item() + print( + f"[grad] {name:55s} rel_max_err={rel:.3e} norm_ratio(orth/base)={ratio:.4f}", + flush=True, + ) + if rel > max_err: + max_err, worst = rel, name + print( + f"[summary] max relative grad error GTP_remat-vs-DP4-baseline = {max_err:.3e} " + f"(worst: {worst})", + flush=True, + ) + assert max_err < 2e-2, ( + f"GTP_remat2xDP2 reduced gradient does not match the no-GTP_remat DP4 baseline " + f"(max rel err {max_err:.3e} on {worst}) -> gtp_remat-axis grad reduce/scaling error." + ) + + +# --------------------------------------------------------------------------- +# Distributed-optimizer + grad-norm path (the production 64-GPU path) +# --------------------------------------------------------------------------- + + +def _build_ddp_distopt_and_optim(stack): + """Real distributed-optimizer setup (Adam), matching the 64-GPU production path.""" + from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig + from megatron.core.optimizer import OptimizerConfig, get_megatron_optimizer + + config = _make_config() + ddp_config = DistributedDataParallelConfig( + use_distributed_optimizer=True, overlap_grad_reduce=False + ) + module = torch.nn.Sequential() + for i, layer in enumerate(stack): + module.add_module(str(i), layer) + ddp_model = DistributedDataParallel(config, ddp_config, module) + opt_config = OptimizerConfig( + optimizer='adam', + lr=0.01, + bf16=True, + use_distributed_optimizer=True, + use_precision_aware_optimizer=False, + main_params_dtype=torch.float32, + main_grads_dtype=torch.float32, + exp_avg_dtype=torch.float32, + exp_avg_sq_dtype=torch.float32, + clip_grad=1.0, # reported grad-norm is computed pre-clip, so this is just for the step + ) + optim = get_megatron_optimizer(opt_config, [ddp_model]) + return ddp_model, optim + + +def _run_step_distopt(ddp_model, optim, rank): + """Mirror production finalize order: finish_grad_sync -> gtp_remat-finalize -> optim.step(). + Returns the optimizer-reported grad-norm (computed pre-clip from the reduced grads).""" + optim.zero_grad() + ddp_model.zero_grad_buffer() + torch.manual_seed(1000 + rank) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + out = x + for layer in ddp_model.module.children(): + out, _ = layer(out, attention_mask=None) + loss = out.float().mean() + loss.backward() + # Production order (finalize_model_grads): reduce across DP first, THEN the gtp_remat finalize. + ddp_model.finish_grad_sync() + from megatron.core.distributed.finalize_model_grads import ( + _allreduce_replicated_grads_over_gtp_remat_group, + ) + + _allreduce_replicated_grads_over_gtp_remat_group([ddp_model]) + _, grad_norm, _ = optim.step() + return float(grad_norm) + + +def _worker_distopt(rank, world_size, port): + from megatron.core import parallel_state as ps + from megatron.core.process_groups_config import ProcessGroupCollection + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + + # ---------- Phase A: baseline, GTP_remat=1 DP=4, dist-opt + Adam ---------- + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=1 + ) + model_parallel_cuda_manual_seed(42) + pgc = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp', 'gtp_remat']) + base_stack = _make_stack(_make_config(), pgc) + for layer in base_stack: + layer.cuda() + for p in base_stack.parameters(): + dist.broadcast(p.data, src=0) + saved = {n: p.data.clone() for n, p in base_stack.named_parameters()} + base_ddp, base_optim = _build_ddp_distopt_and_optim(base_stack) + base_gn = _run_step_distopt(base_ddp, base_optim, rank) + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + + # ---------- Phase B: GTP_remat=2 DP=2, dist-opt + Adam ---------- + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + model_parallel_cuda_manual_seed(42) + pgc = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp', 'gtp_remat']) + gtp_stack = _make_stack(_make_config(), pgc) + for layer in gtp_stack: + layer.cuda() + g = ps.get_gtp_weight_remat_group() + gtp_rank = g.rank() + for name, p in gtp_stack.named_parameters(): + full = saved[name] + if isinstance(p, GTPShardedParam): + ss = p.shape[0] + p.data.copy_(full[gtp_rank * ss : (gtp_rank + 1) * ss]) + else: + p.data.copy_(full) + gtp_ddp, gtp_optim = _build_ddp_distopt_and_optim(gtp_stack) + gtp_gn = _run_step_distopt(gtp_ddp, gtp_optim, rank) + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + + if rank == 0: + ratio = gtp_gn / max(base_gn, 1e-12) + print( + f"\n[distopt grad-norm] baseline={base_gn:.6f} GTP_remat={gtp_gn:.6f} " + f"ratio={ratio:.4f}", + flush=True, + ) + # Same model, same data, gradients proven equal -> grad-norm must match. + torch.testing.assert_close(torch.tensor(gtp_gn), torch.tensor(base_gn), atol=0, rtol=3e-2) + + +# --------------------------------------------------------------------------- +# MoE + EGTP_remat dist-opt grad-norm path (EGTP_remat shards expert weights) +# --------------------------------------------------------------------------- + +NUM_EXPERTS = 4 +MOE_FFN = 256 + + +def _make_moe_config(): + from megatron.core.transformer.transformer_config import TransformerConfig + + return TransformerConfig( + num_attention_heads=NUM_HEADS, + num_layers=NUM_LAYERS, + hidden_size=HIDDEN, + ffn_hidden_size=FFN_HIDDEN, + num_moe_experts=NUM_EXPERTS, + moe_router_topk=2, + moe_ffn_hidden_size=MOE_FFN, + moe_grouped_gemm=True, + moe_token_dispatcher_type="alltoall", + moe_aux_loss_coeff=0.0, + add_bias_linear=False, + params_dtype=dtype, + hidden_dropout=0.0, + attention_dropout=0.0, + bias_dropout_fusion=False, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + ) + + +def _make_moe_stack(config, pg_collection): + from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec + + spec = get_gpt_layer_with_transformer_engine_spec( + num_experts=NUM_EXPERTS, moe_grouped_gemm=True + ) + return torch.nn.ModuleList( + [ + spec.module(config, spec.submodules, layer_number=i + 1, pg_collection=pg_collection) + for i in range(NUM_LAYERS) + ] + ) + + +def _is_expert_param(name, p): + return ('experts' in name) or (not getattr(p, 'allreduce', True)) + + +def _worker_moe_distopt(rank, world_size, port): + from megatron.core import parallel_state as ps + from megatron.core.process_groups_config import ProcessGroupCollection + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + + pgs = ['tp', 'cp', 'gtp_remat', 'ep'] + + # ---------- Phase A: baseline GTP1/EGTP1, EP2 (DP2 dense / expert_dp2) ---------- + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + expert_model_parallel_size=2, + gtp_remat_size=1, + expert_gtp_remat_size=1, + ) + model_parallel_cuda_manual_seed(42) + pgc = ProcessGroupCollection.use_mpu_process_groups(required_pgs=pgs) + base_stack = _make_moe_stack(_make_moe_config(), pgc) + for layer in base_stack: + layer.cuda() + # Broadcast only NON-expert (dense) params; expert weights are EP-local and must + # stay rank-distinct. Save all params per-rank for the GTP_remat phase to mirror. + for name, p in base_stack.named_parameters(): + if not _is_expert_param(name, p): + dist.broadcast(p.data, src=0) + saved = {n: p.data.clone() for n, p in base_stack.named_parameters()} + base_ddp, base_optim = _build_ddp_distopt_and_optim(base_stack) + base_gn = _run_step_distopt(base_ddp, base_optim, rank) + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + + # ---------- Phase B: GTP2/EGTP2, EP2 (EGTP_remat actually shards experts) ---------- + ps.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + expert_model_parallel_size=2, + gtp_remat_size=2, + expert_gtp_remat_size=2, + ) + model_parallel_cuda_manual_seed(42) + pgc = ProcessGroupCollection.use_mpu_process_groups(required_pgs=pgs) + moe_stack = _make_moe_stack(_make_moe_config(), pgc) + for layer in moe_stack: + layer.cuda() + g = ps.get_gtp_weight_remat_group() + eg = ps.get_expert_gtp_weight_remat_group() + gtp_rank, egtp_rank = g.rank(), eg.rank() + n_egtp_sharded = 0 + for name, p in moe_stack.named_parameters(): + full = saved[name] # EP2 layout identical to baseline -> rank-local match + if isinstance(p, GTPShardedParam): + # dense GTP_remat shards over gtp_remat group; expert (EGTP_remat) over egtp_remat. + is_expert = _is_expert_param(name, p) + r = egtp_rank if is_expert else gtp_rank + ss = p.shape[0] + p.data.copy_(full[r * ss : (r + 1) * ss]) + if is_expert: + n_egtp_sharded += 1 + else: + p.data.copy_(full) + if rank == 0: + print( + f"[moe-egtp] egtp-sharded expert params = {n_egtp_sharded} (must be >0 to be a " + f"faithful EGTP_remat test)", + flush=True, + ) + moe_ddp, moe_optim = _build_ddp_distopt_and_optim(moe_stack) + moe_gn = _run_step_distopt(moe_ddp, moe_optim, rank) + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + + if rank == 0: + ratio = moe_gn / max(base_gn, 1e-12) + print( + f"\n[moe distopt grad-norm] baseline={base_gn:.6f} GTP_remat={moe_gn:.6f} " + f"ratio={ratio:.4f}", + flush=True, + ) + torch.testing.assert_close(torch.tensor(moe_gn), torch.tensor(base_gn), atol=0, rtol=3e-2) + + +def _worker_idog_span(rank, world_size, port): + """Dist-opt grad-stats group (intra_dist_opt) must span the FULL world for both + dense-only and MoE(EP2/EGTP2) configs. A naive build collapses the MoE case to a sub-world + group (egtp factored out of expert_data_parallel_size), under-counting the grad-norm.""" + from megatron.core import parallel_state as ps + + # MoE EP2 EGTP2 GTP2 expert config. + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + expert_model_parallel_size=2, + gtp_remat_size=2, + expert_gtp_remat_size=2, + ) + moe_idog = ps.get_intra_distributed_optimizer_instance_group().size() + ps.destroy_model_parallel() + # Dense-only GTP2 (must remain world too). + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + dense_idog = ps.get_intra_distributed_optimizer_instance_group().size() + ps.destroy_model_parallel() + if rank == 0: + print( + f"[idog] MoE intra_dist_opt.size={moe_idog} dense.size={dense_idog} " + f"(world={world_size})", + flush=True, + ) + assert moe_idog == world_size, ( + f"MoE grad-stats group = {moe_idog}, expected world {world_size} " + f"-> grad-norm would under-count gtp_remat/egtp_remat-sharded params" + ) + assert dense_idog == world_size, f"dense grad-stats group = {dense_idog}" + + +class TestGTPGradCorrectness: + def test_distopt_gradstats_group_spans_world(self): + """intra_dist_opt_group (grad-stats) must span the full world.""" + if torch.cuda.device_count() < 4: + pytest.skip("Requires 4 CUDA devices") + _run_distributed(_worker_idog_span, 4) + + @pytest.mark.parametrize("per_token_loss", [False, True]) + def test_gtp2_dp2_grad_matches_dp4_baseline(self, per_token_loss): + """GTP2xDP2 reduced grad must match no-GTP_remat DP4 (non-dist-opt main_grad). + + per_token_loss=True disables DDP's 1/dp pre-scaling and normalizes by + 1/total_global_tokens, so the gtp_remat axis must be SUM-reduced (plain reduce-scatter + + SUM finalize), NOT the 1/gtp MEAN used otherwise. A regression to an unconditional mean + shrinks every gtp grad by 1/gtp and the per_token_loss case catches it (GTP2xDP2 sum-grad + must still match the DP4 sum-grad). GTP_CONFIG is a process-global, so set it for the run + and always reset it. + """ + if torch.cuda.device_count() < 4: + pytest.skip("Requires 4 CUDA devices") + from megatron.core.tensor_parallel.generalized_tensor_parallelism import update_gtp_config + + update_gtp_config(calculate_per_token_loss=per_token_loss) + try: + _run_distributed(_worker, 4, per_token_loss) + finally: + update_gtp_config(calculate_per_token_loss=False) + + def test_gtp2_dp2_distopt_grad_norm_matches_dp4_baseline(self): + """GTP2xDP2 dist-opt grad-norm must match no-GTP_remat DP4 (the 64-GPU path).""" + if torch.cuda.device_count() < 4: + pytest.skip("Requires 4 CUDA devices") + _run_distributed(_worker_distopt, 4) + + @pytest.mark.skip( + reason="EP=2 (engages EGTP_remat) but the minimal test dims (SEQ16 BATCH1 hidden256) hit a " + "token-dispatcher shape error in the alltoall path (RuntimeError shape [2,1,4]). Needs a " + "larger MoE config to run; left as a stub. The real EGTP_remat path is validated at scale " + "(loss matches the GTP1/EGTP1 baseline after the is_gtp/allreduce master-param fix)." + ) + def test_moe_egtp_distopt_grad_norm_matches_baseline(self): + """GTP2/EGTP2 MoE dist-opt grad-norm must match GTP1/EGTP1 baseline (EP=2 both).""" + if torch.cuda.device_count() < 4: + pytest.skip("Requires 4 CUDA devices") + _run_distributed(_worker_moe_distopt, 4) diff --git a/tests/unit_tests/generalized_tensor_parallel/test_gtp_loss_correctness.py b/tests/unit_tests/generalized_tensor_parallel/test_gtp_loss_correctness.py new file mode 100644 index 00000000000..cfc6ff8a492 --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_gtp_loss_correctness.py @@ -0,0 +1,183 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Integration test for GTP correctness. + +Validates that GTP run as a first-class parallelism axis +(world_size = TP * GTP * CP * DP) produces the same per-step loss as a no-GTP +baseline. This is the end-to-end proof that the standalone-GTP rank grid built +in parallel_state trains correctly. + +Mirrors TestAttentionGTPCorrectness. With world=4 and gtp_remat_size=4, GTP +yields dp_replicate=1 and a single shard group [0,1,2,3], so the loss must match +the GTP_remat_size=1 baseline. +""" + +import pytest +import torch +import torch.distributed as dist + +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + +if not HAVE_GTP: + pytest.skip("GTP requires TransformerEngine >= 2.19", allow_module_level=True) + +from transformer_engine.pytorch import fp8_autocast + +from megatron.core.tensor_parallel.generalized_tensor_parallelism import GTPShardedParam +from tests.unit_tests.generalized_tensor_parallel.gtp_test_utils import ( # noqa: F401 (autouse, module-scoped: initializes the dist PG); noqa: F401 (autouse) + _assert_loss_trajectories_match, + _restore_gtp_shards_and_init_main_grad, + _run_distributed, + _torchrun_dist_init, + reset_fp8_state, + reset_gtp_globals, +) + + +def _worker_gtp_loss_correctness(rank, world_size, port): + """Baseline (GTP_remat_size=1, DP=4) vs GTP_remat_size=4 (world=TP1*GTP4*CP1*DP1).""" + from transformer_engine.pytorch.quantization import FP8GlobalStateManager + + from megatron.core import parallel_state as ps + from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec + from megatron.core.process_groups_config import ProcessGroupCollection + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + from megatron.core.transformer.transformer_config import TransformerConfig + + HIDDEN = 4096 + NUM_HEADS = 32 + FFN_HIDDEN = 16384 + NUM_LAYERS = 2 + SEQ = 32 + BATCH = 1 + LR = 0.01 + STEPS = 10 + dtype = torch.bfloat16 + + def make_config(): + return TransformerConfig( + num_attention_heads=NUM_HEADS, + num_layers=NUM_LAYERS, + hidden_size=HIDDEN, + ffn_hidden_size=FFN_HIDDEN, + add_bias_linear=False, + params_dtype=dtype, + hidden_dropout=0.0, + attention_dropout=0.0, + bias_dropout_fusion=False, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + ) + + def make_transformer_stack(config, pg_collection): + spec = get_gpt_layer_with_transformer_engine_spec() + return torch.nn.ModuleList( + [ + spec.module( + config, spec.submodules, layer_number=i + 1, pg_collection=pg_collection + ) + for i in range(NUM_LAYERS) + ] + ) + + def run_step(layers, x): + with fp8_autocast(enabled=False): + for layer in layers: + x, _ = layer(x, attention_mask=None) + return x.mean() + + # ---- Phase 1: Baseline — GTP_remat=1 (DP=4) ---- + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=1 + ) + model_parallel_cuda_manual_seed(42) + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=['tp', 'cp', 'gtp_remat'] + ) + config = make_config() + layers = make_transformer_stack(config, pg_collection) + for layer in layers: + layer.cuda() + for p in layers.parameters(): + dist.broadcast(p.data, src=0) + saved_weights = {n: p.data.clone() for n, p in layers.named_parameters()} + + baseline_losses = [] + for step in range(STEPS): + torch.manual_seed(step) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + dist.broadcast(x, src=0) + loss = run_step(layers, x) + if rank == 0: + baseline_losses.append(loss.item()) + loss.backward() + with torch.no_grad(): + for p in layers.parameters(): + if p.grad is not None: + p.data.sub_(LR * p.grad) + p.grad.zero_() + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + FP8GlobalStateManager.reset() + + # ---- Phase 2: GTP_remat=4 (world = TP1 * GTP4 * CP1 * DP1) ---- + ps.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + gtp_remat_size=4, # standalone-axis GTP_remat under test + ) + model_parallel_cuda_manual_seed(42) + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=['tp', 'cp', 'gtp_remat'] + ) + config = make_config() + layers_gtp = make_transformer_stack(config, pg_collection) + for layer in layers_gtp: + layer.cuda() + + gtp_remat_group = ps.get_gtp_weight_remat_group() + gtp_remat_size = gtp_remat_group.size() + gtp_rank = gtp_remat_group.rank() + assert gtp_remat_size == 4, f"GTP shard group size should be 4, got {gtp_remat_size}" + + gtp_params = [p for p in layers_gtp.parameters() if isinstance(p, GTPShardedParam)] + assert len(gtp_params) > 0, "GTP not active: no GTPShardedParam found" + + _restore_gtp_shards_and_init_main_grad(layers_gtp, saved_weights, gtp_rank, dtype) + + gtp_losses = [] + for step in range(STEPS): + for p in layers_gtp.parameters(): + if isinstance(p, GTPShardedParam): + p.main_grad.zero_() + torch.manual_seed(step) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + dist.broadcast(x, src=0) + loss = run_step(layers_gtp, x) + if rank == 0: + gtp_losses.append(loss.item()) + loss.backward() + with torch.no_grad(): + for p in layers_gtp.parameters(): + if isinstance(p, GTPShardedParam): + p.data.sub_((LR / gtp_remat_size) * p.main_grad) + elif p.grad is not None: + p.data.sub_(LR * p.grad) + p.grad.zero_() + + ps.destroy_model_parallel() + ps.initialize_model_parallel() + GTPShardedParam._chain_state = {} + + if rank == 0: + _assert_loss_trajectories_match(baseline_losses, gtp_losses, STEPS) + + +class TestGTPLossCorrectness: + def test_gtp_loss_trajectory_matches_baseline(self): + """GTP_remat_size=4 per-step losses must match no-GTP baseline (atol=1e-5, rtol=1e-5).""" + if torch.cuda.device_count() < 4: + pytest.skip("Requires at least 4 CUDA devices") + _run_distributed(_worker_gtp_loss_correctness, 4) diff --git a/tests/unit_tests/generalized_tensor_parallel/test_gtp_muon_dcp.py b/tests/unit_tests/generalized_tensor_parallel/test_gtp_muon_dcp.py new file mode 100644 index 00000000000..b26d8a974ce --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_gtp_muon_dcp.py @@ -0,0 +1,327 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Unit tests for GTP + Muon (LayerWise) distributed checkpointing. + +Covers the optimizer-state checkpoint roundtrip for the +:class:`LayerWiseDistributedOptimizer` (Muon) under GTP, where GTP-replicated +matrix params (e.g. the MoE router) are kept whole and must be disambiguated +by ``replica_id`` so DCP does not see multiple writers for the same shard. +""" + +import torch + +from megatron.core.dist_checkpointing import load, save +from tests.unit_tests.dist_checkpointing import TempNamedDir, setup_model_and_optimizer +from tests.unit_tests.test_utilities import Utils + + +def check_equal(input_1, input_2): + """Check if two inputs are equal, used for checking checkpointing.""" + if isinstance(input_1, dict) and isinstance(input_2, dict): + assert input_1.keys() == input_2.keys() + for key in input_1.keys(): + check_equal(input_1[key], input_2[key]) + elif isinstance(input_1, list) and isinstance(input_2, list): + assert len(input_1) == len(input_2) + for i in range(len(input_1)): + check_equal(input_1[i], input_2[i]) + elif isinstance(input_1, torch.Tensor) and isinstance(input_2, torch.Tensor): + assert torch.all(input_1 == input_2), f"Input 1: {input_1} != Input 2: {input_2}" + elif type(input_1) != type(input_2): + assert False, f"Input 1 type: {type(input_1)} != Input 2 type: {type(input_2)}" + else: + assert input_1 == input_2, f"Input 1: {input_1} != Input 2: {input_2}" + + +def _initialize_native_fp8_moe_model( + pre_process=True, + post_process=True, + seed=0, + use_glu=True, + use_sp=False, + use_te=True, + use_grouped_mlp=False, + **config_kwargs, +): + """``initialize_moe_model`` variant for native-FP8 (fp8_model_init / mxfp8) weights. + + ``params_dtype`` is set up front and the ``.bfloat16()`` / ``.random_()`` post-passes skip FP8 + params — a dtype cast or in-place ``random_`` would replace/destroy the native FP8 storage. + """ + from megatron.core.fp8_utils import is_float8tensor + from megatron.core.models.gpt import GPTModel + from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec + from megatron.core.tensor_parallel import model_parallel_cuda_manual_seed + from megatron.core.transformer import TransformerConfig + + # Passed through training.get_model but not part of TransformerConfig. + config_kwargs.pop("pg_collection", None) + config_kwargs.pop("config", None) + torch.manual_seed(seed) + model_parallel_cuda_manual_seed(seed) + expert_num = 8 + + # Dims sized so every GTP4 / EGTP2 FP8 shard dim stays a multiple of the MXFP8 block (32). + default_config_kwargs = dict( + num_layers=2, + hidden_size=128, + num_attention_heads=8, + kv_channels=16, + ffn_hidden_size=256, + use_cpu_initialization=False, + params_dtype=torch.bfloat16, + num_moe_experts=expert_num, + sequence_parallel=use_sp, + moe_grouped_gemm=use_grouped_mlp, + add_bias_linear=False, + fp8='e4m3', + fp8_recipe='mxfp8', + fp8_param=True, + ) + default_config_kwargs.update(**config_kwargs) + transformer_config = TransformerConfig(**default_config_kwargs, gated_linear_unit=use_glu) + spec = get_gpt_layer_with_transformer_engine_spec( + num_experts=expert_num, moe_grouped_gemm=use_grouped_mlp + ) + model = GPTModel( + config=transformer_config, + transformer_layer_spec=spec, + vocab_size=128, + max_sequence_length=4, + pre_process=pre_process, + post_process=post_process, + ) + with torch.no_grad(): + for p in model.parameters(): + if not is_float8tensor(p): + p.random_() + return model + + +class TestGTPMuonDCP: + """GTP + Muon (LayerWise) distributed checkpointing tests.""" + + def teardown_method(self, method): + Utils.destroy_model_parallel() + + def test_gtp_muon_moe_save_load(self, tmp_path_dist_ckpt): + """GTP + Muon (LayerWise) optimizer-state checkpoint roundtrip. + + GTP-REPLICATED, Muon-managed matrix params (e.g. the MoE router, held identically on every + GTP peer) must not collide on GTP peers during checkpoint save: LayerWise keeps each such + param whole, so its optimizer-state ShardedTensor has the same key+offset on all GTP peers + and the replica_id must distinguish them, or DCP validate_sharding_integrity reports 2 + writers ('Invalid access pattern ... [[2]]'). Adam dodges this by sharding the state. + """ + import os + from functools import partial + + import pytest + + from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + + if not HAVE_GTP: + pytest.skip("GTP requires TE with hook registry") + if int(os.environ.get('WORLD_SIZE', '1')) != 4: + pytest.skip("Requires world_size 4 (gtp2 x dp2)") + + os.environ['MEGATRON_GTP_FORCE_ENABLE'] = '1' + from megatron.core import parallel_state as ps + from megatron.core.tensor_parallel import model_parallel_cuda_manual_seed + from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( + GTP_CONFIG, + GTPShardedParam, + update_gtp_config, + ) + from tests.unit_tests.dist_checkpointing.utils import initialize_moe_model + + Utils.initialize_model_parallel(1, 1) # bootstrap torch.distributed + model parallel + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + model_parallel_cuda_manual_seed(2) + # Disable GTP_remat alignment padding so the tiny test dims slice cleanly by gtp_remat_size. + _orig_pad = GTP_CONFIG.pad_for_alignment + update_gtp_config(pad_for_alignment=0) + # GTP_remat dims (divisible by gtp_remat_size=2); GPU init (CPU affine not GTP_remat-aware + # for the strided QKV weight). + moe_cfg = dict( + hidden_size=64, + num_attention_heads=8, + kv_channels=8, + ffn_hidden_size=128, + use_cpu_initialization=False, + ) + meta = {'distrib_optim_sharding_type': 'dp_reshardable'} + with TempNamedDir(tmp_path_dist_ckpt / 'gtp_muon_moe_A', sync=True) as ckpt_dir_A: + with TempNamedDir(tmp_path_dist_ckpt / 'gtp_muon_moe_B', sync=True) as ckpt_dir_B: + model_A, optimizer_A = setup_model_and_optimizer( + seed=2, + tp=1, + pp=1, + bf16=True, + dist_opt=True, + use_param_layout=True, + initialize_fn=partial(initialize_moe_model, use_te=True, **moe_cfg), + optimizer='dist_muon', + ) + assert any( + isinstance(p, GTPShardedParam) for p in model_A[0].parameters() + ), "GTP not active: no GTPShardedParam in the GTP_remat_size=2 MoE model" + + model_sd_A = model_A[0].sharded_state_dict() + optim_sd_A = optimizer_A.sharded_state_dict(model_sd_A, metadata=meta) + save( + optim_sd_A, ckpt_dir_A + ) # fails (2 writers) before the LayerWise replica_id fix + + model_B, optimizer_B = setup_model_and_optimizer( + seed=3, + tp=1, + pp=1, + bf16=True, + dist_opt=True, + use_param_layout=True, + initialize_fn=partial(initialize_moe_model, use_te=True, **moe_cfg), + optimizer='dist_muon', + ) + model_sd_B = model_B[0].sharded_state_dict() + load_sharded_sd = optimizer_B.sharded_state_dict( + model_sd_B, is_loading=True, metadata=meta + ) + state_dict = load(load_sharded_sd, ckpt_dir_A) + optimizer_B.load_state_dict(state_dict) + optim_sd_B = optimizer_B.sharded_state_dict(model_sd_B, metadata=meta) + save(optim_sd_B, ckpt_dir_B) + + update_gtp_config(pad_for_alignment=_orig_pad) + + Utils.destroy_model_parallel() + Utils.initialize_model_parallel(1, 1) + from megatron.core.dist_checkpointing import load_plain_tensors + + check_equal(load_plain_tensors(ckpt_dir_A), load_plain_tensors(ckpt_dir_B)) + + def test_gtp_muon_moe_native_fp8_save_load(self, tmp_path_dist_ckpt): + """GTP + Muon (LayerWise) + native-FP8 (fp8_model_init / mxfp8) checkpoint roundtrip. + + Regression guard for the native-FP8 optimizer-state save. Native-FP8 GTP weights are + dequantized into a NEW bf16 tensor when the model builds its ShardedTensor + (make_tp_sharded_tensor_for_checkpoint), which breaks the ``id(entry.data) == id(param)`` + match every native-FP8 GTP param relies on -- so ALL of them fall into + ``_backfill_gtp_sharded_param_map``. The fix reuses each model entry (via the + ``_gtp_dequant_src`` backlink), preserving its offsets/replica_id; this exercises that + reuse path end-to-end and asserts a bit-exact save/load/re-save roundtrip. + + Uses the gtp_remat-only MoE grid (no expert-parallel), matching the sibling bf16 test. + The specific [[2],[2]] cross-expert collision from the production crash needs the GROUPED / + EGTP expert grid (shared key + per-expert offset, where the old EP-unaware rebuild dropped + that offset). That grid is intentionally avoided here: LayerWiseDistributedOptimizer's + whole-param LPT layout over EGTP-sharded expert weights leaves a coverage hole that fails + the same save even in pure BF16 (independent of native-FP8 and of this fix). + """ + import os + from functools import partial + from unittest import mock + + import pytest + + from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + from tests.unit_tests.dist_checkpointing import utils as _dc_utils + from tests.unit_tests.generalized_tensor_parallel.gtp_test_utils import _requires_mxfp8 + + if not HAVE_GTP: + pytest.skip("GTP requires TE with hook registry") + if int(os.environ.get('WORLD_SIZE', '1')) != 4: + pytest.skip("Requires world_size 4 (gtp2 x dp2)") + _requires_mxfp8() + + os.environ['MEGATRON_GTP_FORCE_ENABLE'] = '1' + from megatron.core import parallel_state as ps + from megatron.core.fp8_utils import is_float8tensor + from megatron.core.tensor_parallel import model_parallel_cuda_manual_seed + from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( + GTP_CONFIG, + is_gtp_param, + tag_gtp_params_with_names, + update_gtp_config, + ) + + Utils.initialize_model_parallel(1, 1) # bootstrap torch.distributed + model parallel + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + model_parallel_cuda_manual_seed(2) + # MXFP8 needs shard dims % 32; padding off so dims slice cleanly by the remat size. + _orig_pad = GTP_CONFIG.pad_for_alignment + update_gtp_config(pad_for_alignment=0) + + # MXFP8 params can't be aliased into the DDP param buffer (replace_raw_data unsupported); + # the production native-FP8 path sets reuse_grad_buf_for_mxfp8_param_ag to skip that + # aliasing. The shared harness builds its own mock args, so wrap init_basic_mock_args to + # flip the flags before the DDP config is built from them. + _orig_init_args = _dc_utils.init_basic_mock_args + + def _init_args_fp8(args, tp, pp, bf16=True): + _orig_init_args(args, tp, pp, bf16=bf16) + args.fp8_param_gather = True + args.reuse_grad_buf_for_mxfp8_param_ag = True + return args + + init_fn = partial(_initialize_native_fp8_moe_model, use_te=True, use_grouped_mlp=False) + meta = {'distrib_optim_sharding_type': 'dp_reshardable'} + with ( + mock.patch.object(_dc_utils, 'init_basic_mock_args', _init_args_fp8), + TempNamedDir(tmp_path_dist_ckpt / 'gtp_muon_fp8_A', sync=True) as ckpt_dir_A, + TempNamedDir(tmp_path_dist_ckpt / 'gtp_muon_fp8_B', sync=True) as ckpt_dir_B, + ): + model_A, optimizer_A = setup_model_and_optimizer( + seed=2, + tp=1, + pp=1, + bf16=True, + dist_opt=True, + use_param_layout=True, + initialize_fn=init_fn, + optimizer='dist_muon', + ) + tag_gtp_params_with_names(model_A[0]) + assert any( + is_gtp_param(p) and is_float8tensor(p) for p in model_A[0].parameters() + ), "no native-FP8 GTP param present; test is not exercising the FP8 path" + + model_sd_A = model_A[0].sharded_state_dict() + # Every native-FP8 GTP param is unmatched (dequantized copy) and reuses its model + # entry via _backfill_gtp_sharded_param_map; save validates the composed sharding. + optim_sd_A = optimizer_A.sharded_state_dict(model_sd_A, metadata=meta) + save(optim_sd_A, ckpt_dir_A) + + model_B, optimizer_B = setup_model_and_optimizer( + seed=3, + tp=1, + pp=1, + bf16=True, + dist_opt=True, + use_param_layout=True, + initialize_fn=init_fn, + optimizer='dist_muon', + ) + tag_gtp_params_with_names(model_B[0]) + model_sd_B = model_B[0].sharded_state_dict() + load_sharded_sd = optimizer_B.sharded_state_dict( + model_sd_B, is_loading=True, metadata=meta + ) + state_dict = load(load_sharded_sd, ckpt_dir_A) + optimizer_B.load_state_dict(state_dict) + optim_sd_B = optimizer_B.sharded_state_dict(model_sd_B, metadata=meta) + save(optim_sd_B, ckpt_dir_B) + + update_gtp_config(pad_for_alignment=_orig_pad) + + Utils.destroy_model_parallel() + Utils.initialize_model_parallel(1, 1) + from megatron.core.dist_checkpointing import load_plain_tensors + + check_equal(load_plain_tensors(ckpt_dir_A), load_plain_tensors(ckpt_dir_B)) diff --git a/tests/unit_tests/generalized_tensor_parallel/test_mamba_gtp.py b/tests/unit_tests/generalized_tensor_parallel/test_mamba_gtp.py new file mode 100644 index 00000000000..152b551e78d --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_mamba_gtp.py @@ -0,0 +1,343 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Integration tests for GTP + Mamba correctness. + +Test groups +----------- +TestMambaGTPCorrectness - GTP Mamba loss trajectory matches baseline (no-GTP) over 10 + training steps using MXFP8 and Nemotron3-Super Mamba hyperparameters. +""" + +import pytest +import torch +import torch.distributed as dist + +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + +if not HAVE_GTP: + pytest.skip("GTP requires TransformerEngine >= 2.19", allow_module_level=True) + +from transformer_engine.pytorch import fp8_autocast, fp8_model_init + +from megatron.core.tensor_parallel.generalized_tensor_parallelism import GTPShardedParam +from tests.unit_tests.generalized_tensor_parallel.gtp_test_utils import ( + _assert_loss_trajectories_match, + _requires_mxfp8, + _restore_gtp_shards_and_init_main_grad, + _run_distributed, + _torchrun_dist_init, + reset_fp8_state, + reset_gtp_globals, +) + +# --------------------------------------------------------------------------- +# Mamba GTP_remat correctness: per-step loss trajectory baseline vs GTP_remat=4 +# --------------------------------------------------------------------------- + + +def _worker_mamba_gtp_correctness(rank, world_size, port): + """Verify GTP Mamba produces the same per-step loss as a no-GTP baseline. + + Phase 1 — GTP_remat_size=1, DP=4: + All 4 ranks hold the full model and process identical inputs. Gradients + are identical across ranks (no all-reduce needed). Weight update: + param.data -= lr * param.grad + + Phase 2 — GTP_remat_size=4, DP=1: + Weights sharded across 4 ranks. After backward, wgrad reduce-scatter + sums each shard's identical wgrad over all ranks, so: + main_grad[rank_i] = gtp_remat_size * dW[shard_i] + The optimizer divides by gtp_remat_size to recover the per-element gradient: + param.data -= (lr / gtp_remat_size) * param.main_grad + + Both phases use identical initial weights (synced from rank 0 in phase 1, + restored as shards in phase 2) and identical step-by-step inputs. The + per-step loss trajectories must agree within 0.1% relative error. + """ + from transformer_engine.common.recipe import MXFP8BlockScaling + from transformer_engine.pytorch.quantization import FP8GlobalStateManager + + from megatron.core import parallel_state as ps + from megatron.core.extensions.transformer_engine import ( + TELayerNormColumnParallelLinear, + TERowParallelLinear, + ) + from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add + from megatron.core.process_groups_config import ProcessGroupCollection + from megatron.core.ssm.mamba_layer import MambaLayer, MambaLayerSubmodules + from megatron.core.ssm.mamba_mixer import MambaMixer, MambaMixerSubmodules + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + from megatron.core.transformer.spec_utils import ModuleSpec + from megatron.core.transformer.transformer_config import TransformerConfig + + # Nemotron3-Super Proxy Mamba hyperparameters. + # in_proj_out = 2*8192 + 2*8*128 + 128 = 18560; 18560/4 = 4640, 4640%16 = 0 (MXFP8-aligned). + HIDDEN = 4096 + NHEADS = 128 # mamba_num_heads; d_inner = nheads * headdim = 128 * 64 = 8192 + NGROUPS = 8 # mamba_num_groups (default) + D_STATE = 128 # mamba_state_dim (default) + NUM_LAYERS = 2 + SEQ = 32 + BATCH = 1 + LR = 0.01 + STEPS = 10 + dtype = torch.bfloat16 + recipe = MXFP8BlockScaling() # native-FP8 Phase 3 (Phases 1-2 run in BF16) + + def make_config(): + return TransformerConfig( + num_attention_heads=32, + num_layers=NUM_LAYERS, + hidden_size=HIDDEN, + mamba_num_heads=NHEADS, + mamba_head_dim=64, + mamba_state_dim=D_STATE, + mamba_num_groups=NGROUPS, + use_mamba_mem_eff_path=True, + params_dtype=dtype, + hidden_dropout=0.0, + bias_dropout_fusion=False, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + ) + + def make_mamba_stack(config, pg_collection): + submodules = MambaLayerSubmodules( + mixer=ModuleSpec( + module=MambaMixer, + submodules=MambaMixerSubmodules( + in_proj=TELayerNormColumnParallelLinear, out_proj=TERowParallelLinear + ), + ), + mamba_bda=get_bias_dropout_add, + ) + return torch.nn.ModuleList( + [ + MambaLayer(config, submodules, layer_number=i + 1, pg_collection=pg_collection) + for i in range(NUM_LAYERS) + ] + ) + + def run_step(layers, x): + with fp8_autocast(enabled=False): + for layer in layers: + x = layer(x) + return x.mean() + + # ------------------------------------------------------------------------- + # Phase 1: Baseline — GTP_remat=1 (DP=4) + # ------------------------------------------------------------------------- + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=1 + ) + model_parallel_cuda_manual_seed(42) + + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=['tp', 'cp', 'gtp_remat'] + ) + config = make_config() + layers = make_mamba_stack(config, pg_collection) + for layer in layers: + layer.cuda() + + # Verify baseline has no GTP_remat sharding (gtp_remat_size=1 should leave plain parameters). + assert not any( + isinstance(p, GTPShardedParam) for p in layers.parameters() + ), "Baseline GTP_remat_size=1 stack should have no GTPShardedParam" + + # Synchronize weights from rank 0 across all DP ranks. + for p in layers.parameters(): + dist.broadcast(p.data, src=0) + + # Save initial weights; will be used to initialize the GTP_remat model identically. + saved_weights = {n: p.data.clone() for n, p in layers.named_parameters()} + + baseline_losses = [] + for step in range(STEPS): + torch.manual_seed(step) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + dist.broadcast(x, src=0) + + loss = run_step(layers, x) + if rank == 0: + baseline_losses.append(loss.item()) + + loss.backward() + with torch.no_grad(): + for p in layers.parameters(): + if p.grad is not None: + p.data.sub_(LR * p.grad) + p.grad.zero_() + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + FP8GlobalStateManager.reset() + + # ------------------------------------------------------------------------- + # Phase 2: GTP_remat=4 (DP=1) + # ------------------------------------------------------------------------- + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=4 + ) + model_parallel_cuda_manual_seed(42) + + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=['tp', 'cp', 'gtp_remat'] + ) + config = make_config() + layers_gtp = make_mamba_stack(config, pg_collection) + for layer in layers_gtp: + layer.cuda() + + gtp_remat_group = ps.get_gtp_weight_remat_group() + gtp_remat_size = gtp_remat_group.size() + gtp_rank = gtp_remat_group.rank() + + # Verify GTP_remat is truly active: at least one param must be a GTPShardedParam. + gtp_params = [p for p in layers_gtp.parameters() if isinstance(p, GTPShardedParam)] + assert ( + len(gtp_params) > 0 + ), "GTP is not active: no GTPShardedParam found in GTP_remat_size=4 Mamba stack" + + # Restore initial weights into shards and pre-allocate main_grad for the backward. + _restore_gtp_shards_and_init_main_grad(layers_gtp, saved_weights, gtp_rank, dtype) + + gtp_losses = [] + for step in range(STEPS): + for p in layers_gtp.parameters(): + if isinstance(p, GTPShardedParam): + p.main_grad.zero_() + + torch.manual_seed(step) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + dist.broadcast(x, src=0) + + loss = run_step(layers_gtp, x) + if rank == 0: + gtp_losses.append(loss.item()) + + loss.backward() + + # After RS, main_grad = gtp_remat_size * dW_shard (sum over ranks, all ranks hold the same + # full wgrad after all-gathering the weight in fwd). Divide by gtp_remat_size so the weight + # update is equivalent to the baseline. + with torch.no_grad(): + for p in layers_gtp.parameters(): + if isinstance(p, GTPShardedParam): + p.data.sub_((LR / gtp_remat_size) * p.main_grad) + elif p.grad is not None: + p.data.sub_(LR * p.grad) + p.grad.zero_() + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + FP8GlobalStateManager.reset() + + # ------------------------------------------------------------------------- + # Phase 3: GTP_remat=4 (DP=1), NATIVE MXFP8 weights (--fp8-param-gather path). + # in_proj/out_proj are built under fp8_model_init -> native MXFP8 GTP shards. FP8 params can't + # be updated in place, so this leg keeps an FP32 master and re-quantizes each step via + # gtp_native_fp8_load_context (same copy_ mechanism as checkpoint load). Loss must track the + # BF16 baseline within MXFP8 noise; a frozen/miswired FP8 weight flattens or diverges it. + # ------------------------------------------------------------------------- + from megatron.core.fp8_utils import is_float8tensor + from megatron.core.tensor_parallel.gtp_api import gtp_native_fp8_load_context, is_gtp_param + + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=4 + ) + model_parallel_cuda_manual_seed(42) + pg_collection = ProcessGroupCollection.use_mpu_process_groups( + required_pgs=['tp', 'cp', 'gtp_remat'] + ) + config = make_config() + with fp8_model_init(enabled=True, recipe=recipe): + layers_fp8 = make_mamba_stack(config, pg_collection) + for layer in layers_fp8: + layer.cuda() + + # Verify native-FP8 GTP is truly active: some weight must be a native FP8 GTP shard. + native_fp8 = [p for p in layers_fp8.parameters() if is_gtp_param(p) and is_float8tensor(p)] + assert len(native_fp8) > 0, "No native-FP8 GTP weight found in fp8_model_init mamba stack" + + # Init: FP32 master per param = the gtp_rank shard of the saved baseline weights (padded to the + # native-FP8 shard size); FP8 params re-quantized from their master via the load context. + masters = {} + with torch.no_grad(), gtp_native_fp8_load_context(layers_fp8): + for name, p in layers_fp8.named_parameters(): + full = saved_weights[name] + if is_gtp_param(p): + shard = p.shape[0] # native-FP8 shard may include GTP alignment pad rows + aligned = shard * gtp_remat_size + if full.shape[0] < aligned: + full = torch.nn.functional.pad(full, (0, 0, 0, aligned - full.shape[0])) + m = full[gtp_rank * shard : (gtp_rank + 1) * shard].float().clone() + else: + m = full.float().clone() + masters[name] = m + p.copy_(m.to(dtype)) # BF16->FP8 (inside the load context) or plain BF16 copy + for p in layers_fp8.parameters(): + if is_gtp_param(p): + p.main_grad = torch.zeros(p.shape, dtype=dtype, device='cuda') + + fp8_losses = [] + for step in range(STEPS): + for p in layers_fp8.parameters(): + if is_gtp_param(p): + p.main_grad.zero_() + + torch.manual_seed(step) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + dist.broadcast(x, src=0) + + with fp8_autocast(enabled=True, fp8_recipe=recipe): + y = x + for layer in layers_fp8: + y = layer(y) + loss = y.mean() + if rank == 0: + fp8_losses.append(loss.item()) + loss.backward() + + # Update FP32 masters (same math as bf16 Phase 2), then re-quantize FP8 shards from them. + with torch.no_grad(): + for name, p in layers_fp8.named_parameters(): + if is_gtp_param(p): + masters[name].sub_((LR / gtp_remat_size) * p.main_grad.float()) + elif p.grad is not None: + masters[name].sub_(LR * p.grad.float()) + p.grad.zero_() + with gtp_native_fp8_load_context(layers_fp8): + for name, p in layers_fp8.named_parameters(): + p.copy_(masters[name].to(dtype)) + + ps.destroy_model_parallel() + ps.initialize_model_parallel() + GTPShardedParam._chain_state = {} + FP8GlobalStateManager.reset() + + # ------------------------------------------------------------------------- + # Compare per-step loss trajectories on rank 0 + # ------------------------------------------------------------------------- + if rank == 0: + _assert_loss_trajectories_match(baseline_losses, gtp_losses, STEPS) + # Native-FP8 leg tracks the baseline within MXFP8 noise (looser than the bf16-vs-bf16 tol). + import torch as _torch + + diff = (_torch.tensor(fp8_losses) - _torch.tensor(baseline_losses)).abs().max().item() + assert diff < 0.2, ( + f"Native-FP8 GTP mamba loss diverges from BF16 baseline " + f"(max per-step |diff|={diff:.4f}; fp8: {fp8_losses[0]:.3f}->{fp8_losses[-1]:.3f}, " + f"bf16: {baseline_losses[0]:.3f}->{baseline_losses[-1]:.3f})." + ) + + +class TestMambaGTPCorrectness: + def test_mamba_gtp_loss_trajectory_matches_baseline(self): + """GTP Mamba per-step losses must match the no-GTP baseline for BOTH bf16 (Phase 2) and + native-MXFP8 (Phase 3) weight-remat, within bf16 reduction / mxfp8 quantization noise.""" + if torch.cuda.device_count() < 4: + pytest.skip("Requires at least 4 CUDA devices") + _requires_mxfp8() # Phase 3 builds native-FP8 mamba weights (fp8_model_init) + _run_distributed(_worker_mamba_gtp_correctness, 4) diff --git a/tests/unit_tests/generalized_tensor_parallel/test_moe_egtp.py b/tests/unit_tests/generalized_tensor_parallel/test_moe_egtp.py new file mode 100644 index 00000000000..da4e957fccd --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_moe_egtp.py @@ -0,0 +1,342 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Integration tests for EGTP_remat + MoE correctness. + +Test groups +----------- +TestMoEEGTPCorrectness - EGTP_remat MoE loss trajectory matches baseline (no-EGTP_remat) over 10 + training steps using MXFP8 and Nemotron3-Super MoE hyperparameters. +""" + +import pytest +import torch +import torch.distributed as dist + +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + +if not HAVE_GTP: + pytest.skip("GTP requires TransformerEngine >= 2.19", allow_module_level=True) + +from transformer_engine.pytorch import fp8_autocast + +from megatron.core.tensor_parallel.generalized_tensor_parallelism import GTPShardedParam +from megatron.core.transformer.moe.moe_utils import get_default_pg_collection +from tests.unit_tests.generalized_tensor_parallel.gtp_test_utils import ( + _assert_loss_trajectories_match, + _run_distributed, + _torchrun_dist_init, + reset_fp8_state, + reset_gtp_globals, +) + +# --------------------------------------------------------------------------- +# MoE EGTP_remat correctness: per-step loss trajectory EP=4 baseline vs EP=2+EGTP_remat=2 +# --------------------------------------------------------------------------- + + +def _worker_moe_egtp_correctness(rank, world_size, port): + """Verify EP=2+EGTP_remat=2 MoE matches per-step loss of EP=4 no-EGTP_remat baseline. + + Phase 1 — EP=4, EGTP_remat=1: + All 4 ranks form one EP group; each rank holds 2 full expert weights (8 total). + All ranks receive the same MoE-layer input; alltoall dispatch routes each token + to its assigned expert rank, so each rank computes a different token subset. + Gradients are local to each expert's rank. Weight update: + param.data -= lr * param.grad + + Phase 2 — EP=2, EGTP_remat=2: + Two EP groups of 2 ranks, each EGTP_remat-sharded over 2 ranks. Expert weights + are sharded along dim 0 within each EGTP_remat group (shard = full_dim0 / egtp_remat_size). + After backward, wgrad reduce-scatter sums each shard's identical wgrad: + main_grad[rank_i] = egtp_remat_size * dW[shard_i] + The optimizer divides by egtp_remat_size: + param.data -= (lr / egtp_remat_size) * param.main_grad + + Weight sharing (test-only): + To ensure both phases start from identical expert weights, an all-gather + collects the full 8-expert table from the EP=4 group (where each rank holds + only 2 experts) onto every rank. Phase 2 then slices each rank's local + experts and EGTP_remat shard from that global table. + + Nemotron3-Super Proxy MoE hyperparameters (scaled for unit-test speed): + hidden=4096, ffn_hidden_size=2688, num_experts=8, topk=2 + MXFP8 alignment with EGTP_remat=2: + 2688/2=1344, 1344%16=0 (fc1 shard); 4096/2=2048, 2048%16=0 (fc2 shard) + """ + from transformer_engine.pytorch.quantization import FP8GlobalStateManager + + from megatron.core import parallel_state as ps + from megatron.core.models.gpt.moe_module_specs import get_moe_module_spec + from megatron.core.process_groups_config import ProcessGroupCollection + from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed + from megatron.core.transformer.transformer_config import TransformerConfig + + # Nemotron3-Super MoE hyperparameters (num_experts scaled from 512 to 8 for test speed). + HIDDEN = 4096 + FFN_HIDDEN = 2688 + NUM_EXPERTS = 8 + TOPK = 2 + SEQ = 32 + BATCH = 1 + LR = 0.01 + STEPS = 10 + dtype = torch.bfloat16 + + def make_config(): + return TransformerConfig( + num_attention_heads=32, + num_layers=1, + hidden_size=HIDDEN, + num_moe_experts=NUM_EXPERTS, + moe_router_topk=TOPK, + moe_ffn_hidden_size=FFN_HIDDEN, + moe_grouped_gemm=True, + moe_token_dispatcher_type="alltoall", + moe_aux_loss_coeff=0.0, + add_bias_linear=False, + params_dtype=dtype, + hidden_dropout=0.0, + bias_dropout_fusion=False, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + ) + + def make_moe_layer(config, pg_collection): + moe_spec = get_moe_module_spec(use_te=True, num_experts=NUM_EXPERTS, moe_grouped_gemm=True) + return moe_spec(config, layer_number=1, pg_collection=pg_collection) + + def run_step(layer, x): + with fp8_autocast(enabled=False): + output, _ = layer(x) + return output.mean() + + # ------------------------------------------------------------------------- + # Phase 1: Baseline — EP=4, EGTP_remat=1 (DP=1) + # ------------------------------------------------------------------------- + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + expert_model_parallel_size=4, + expert_gtp_remat_size=1, + ) + model_parallel_cuda_manual_seed(42) + + pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['ep']) + ep_group = pg_collection.ep + num_local_experts_baseline = NUM_EXPERTS // 4 # = 2 + + config = make_config() + layer = make_moe_layer(config, None) # MoELayer uses get_default_pg_collection() + layer.cuda() + + # Verify baseline has no GTP_remat sharding (EGTP_remat=1 should leave plain parameters). + assert not any( + isinstance(p, GTPShardedParam) for p in layer.parameters() + ), "Baseline EP=4 layer should have no GTPShardedParam (EGTP_remat=1)" + + # Synchronize non-expert weights from rank 0; expert weights are rank-local. + for name, p in layer.named_parameters(): + if 'linear_fc1.weight' not in name and 'linear_fc2.weight' not in name: + dist.broadcast(p.data, src=0) + + # Collect the full expert weight table so Phase 2 can restore identical init weights. + # EP=4: each rank holds 2 experts; all-gather gives every rank the complete [8, dim, ...] table. + local_fc1 = torch.stack( + [ + dict(layer.named_parameters())[f'experts.linear_fc1.weight{i}'].data + for i in range(num_local_experts_baseline) + ] + ) # [2, FFN_HIDDEN, HIDDEN] + global_fc1 = torch.zeros(NUM_EXPERTS, FFN_HIDDEN, HIDDEN, dtype=dtype, device='cuda') + dist.all_gather_into_tensor(global_fc1, local_fc1, group=ep_group) + + local_fc2 = torch.stack( + [ + dict(layer.named_parameters())[f'experts.linear_fc2.weight{i}'].data + for i in range(num_local_experts_baseline) + ] + ) # [2, HIDDEN, FFN_HIDDEN] + global_fc2 = torch.zeros(NUM_EXPERTS, HIDDEN, FFN_HIDDEN, dtype=dtype, device='cuda') + dist.all_gather_into_tensor(global_fc2, local_fc2, group=ep_group) + + # Save non-expert param values (router, norms, etc.) from rank 0. + non_expert_weights = {} + for name, p in layer.named_parameters(): + if 'linear_fc1.weight' not in name and 'linear_fc2.weight' not in name: + non_expert_weights[name] = p.data.clone() + + baseline_losses = [] + for step in range(STEPS): + torch.manual_seed(step) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + dist.broadcast(x, src=0) + + loss = run_step(layer, x) + if rank == 0: + baseline_losses.append(loss.item()) + + loss.backward() + with torch.no_grad(): + for p in layer.parameters(): + if p.grad is not None: + p.data.sub_(LR * p.grad) + p.grad.zero_() + + ps.destroy_model_parallel() + GTPShardedParam._chain_state = {} + FP8GlobalStateManager.reset() + + # ------------------------------------------------------------------------- + # Phase 2: EP=2, EGTP_remat=2 (DP=1 effective) + # ------------------------------------------------------------------------- + ps.initialize_model_parallel( + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + expert_model_parallel_size=2, + expert_gtp_remat_size=2, + ) + model_parallel_cuda_manual_seed(42) + + pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['expt_gtp_remat']) + egtp_remat_group = pg_collection.expt_gtp_remat + egtp_remat_size = egtp_remat_group.size() + egtp_rank = egtp_remat_group.rank() + ep_rank_egtp = dist.get_rank(ps.get_expert_model_parallel_group()) + num_local_experts_egtp = NUM_EXPERTS // 2 # = 4 + + config = make_config() + # Build full pg_collection for MoELayer: default groups + expt_gtp for EGTP_remat sharding. + moe_pg = get_default_pg_collection() + moe_pg.expt_gtp_remat = egtp_remat_group + layer_egtp = make_moe_layer(config, moe_pg) + layer_egtp.cuda() + + # Verify EGTP_remat is truly active: expert weight params must be GTPShardedParam instances. + egtp_params = [p for p in layer_egtp.parameters() if isinstance(p, GTPShardedParam)] + assert len(egtp_params) > 0, "EGTP_remat inactive: no GTPShardedParam in EP=2+EGTP_remat=2" + + # Restore weights from saved global tables. + # Expert local index j → global expert id = ep_rank_egtp * num_local_experts_egtp + j. + fc1_shard = FFN_HIDDEN // egtp_remat_size # 2688/2 = 1344 + fc2_shard = HIDDEN // egtp_remat_size # 4096/2 = 2048 + for name, p in layer_egtp.named_parameters(): + if 'linear_fc1.weight' in name: + j = int(name.rsplit('weight', 1)[1]) + gid = ep_rank_egtp * num_local_experts_egtp + j + p.data.copy_(global_fc1[gid, egtp_rank * fc1_shard : (egtp_rank + 1) * fc1_shard]) + elif 'linear_fc2.weight' in name: + j = int(name.rsplit('weight', 1)[1]) + gid = ep_rank_egtp * num_local_experts_egtp + j + p.data.copy_(global_fc2[gid, egtp_rank * fc2_shard : (egtp_rank + 1) * fc2_shard]) + elif name in non_expert_weights: + p.data.copy_(non_expert_weights[name]) + + # Pre-allocate main_grad for EGTP_remat params (required before the first backward). + for p in layer_egtp.parameters(): + if isinstance(p, GTPShardedParam): + p.main_grad = torch.zeros(p.shape, dtype=dtype, device='cuda') + + egtp_losses = [] + for step in range(STEPS): + for p in layer_egtp.parameters(): + if isinstance(p, GTPShardedParam): + p.main_grad.zero_() + + torch.manual_seed(step) + x = torch.randn(SEQ, BATCH, HIDDEN, dtype=dtype, device='cuda') + dist.broadcast(x, src=0) + + loss = run_step(layer_egtp, x) + if rank == 0: + egtp_losses.append(loss.item()) + + loss.backward() + + # After RS, main_grad = egtp_remat_size * dW_shard. Divide by egtp_remat_size for baseline. + with torch.no_grad(): + for p in layer_egtp.parameters(): + if isinstance(p, GTPShardedParam): + p.data.sub_((LR / egtp_remat_size) * p.main_grad) + elif p.grad is not None: + p.data.sub_(LR * p.grad) + p.grad.zero_() + + ps.destroy_model_parallel() + ps.initialize_model_parallel() + GTPShardedParam._chain_state = {} + + # ------------------------------------------------------------------------- + # Compare per-step loss trajectories on rank 0 + # ------------------------------------------------------------------------- + if rank == 0: + _assert_loss_trajectories_match(baseline_losses, egtp_losses, STEPS, label="egtp_remat") + + +def _worker_expert_bias_gtp_inclusive(rank, world_size, port): + """Router expert-bias must stay identical across gtp_remat peers. + + The aux-loss-free balancer (``get_updated_expert_bias``) all-reduces per-expert token counts + over the router's tp_dp_cp group, then sign-updates the replicated expert_bias. gtp_remat peers + hold DISTINCT tokens but share the replicated router, so the reduction must span gtp_remat. + ``get_tensor_and_data_parallel_group`` therefore spans gtp_remat (like dp); over a gtp-EXCLUDED + group the peers reduce different token sums and the bias diverges -> routing instability + (loss spikes). world=4 -> tp1 * gtp_remat2 * dp2. + """ + import torch.distributed as dist + + from megatron.core import parallel_state as ps + from megatron.core.transformer.moe.moe_utils import get_updated_expert_bias + from megatron.core.utils import get_pg_size + + ps.destroy_model_parallel() + ps.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1, gtp_remat_size=2 + ) + gtp_group = ps.get_gtp_weight_remat_group() + tp_dp_cp = ps.get_tensor_and_data_parallel_group(with_context_parallel=True) # spans gtp_remat + # gtp-EXCLUDED replicate group (explicit, since get_data_parallel_group now defaults inclusive). + replicate = ps.get_data_parallel_group(with_context_parallel=True, with_gtp_remat=False) + + num_experts = 8 + torch.manual_seed(1000 + rank) # distinct per-rank tokens (distinct data on each rank) + base = torch.randint(0, 100, (num_experts,), device="cuda").float() + + def _max_bias_diff_across_gtp(group): + bias = torch.zeros(num_experts, device="cuda") # identical start on every rank + updated = get_updated_expert_bias(base.clone(), bias, 0.01, tp_dp_cp_group=group) + buf = [torch.zeros_like(updated) for _ in range(gtp_group.size())] + dist.all_gather(buf, updated, group=gtp_group) + return max((buf[i] - buf[0]).abs().max().item() for i in range(gtp_group.size())) + + # tp_dp_cp must span gtp_remat (size = replicate_dp_cp x gtp_remat at tp=1). + spans_gtp = get_pg_size(tp_dp_cp) == get_pg_size(replicate) * get_pg_size(gtp_group) + diff_excluded = _max_bias_diff_across_gtp(replicate) + diff_included = _max_bias_diff_across_gtp(tp_dp_cp) + + ps.destroy_model_parallel() + ps.initialize_model_parallel() + + if rank == 0: + assert spans_gtp, "tp_dp_cp group must span the gtp_remat axis (like dp)" + # gtp-excluded reduction reproduces the bug; tp_dp_cp (spans gtp_remat) keeps peers in sync. + assert ( + diff_excluded > 0 + ), f"gtp-excluded group should diverge across gtp_remat peers, got {diff_excluded}" + assert ( + diff_included == 0 + ), f"tp_dp_cp must keep expert_bias identical across gtp_remat peers, got {diff_included}" + + +class TestMoEEGTPCorrectness: + def test_moe_egtp_loss_trajectory_matches_baseline(self): + """EP=2+EGTP_remat=2 MoE per-step losses match EP=4 baseline: atol=rtol=1e-5; MXFP8""" + if torch.cuda.device_count() < 4: + pytest.skip("Requires at least 4 CUDA devices") + _run_distributed(_worker_moe_egtp_correctness, 4) + + def test_expert_bias_gtp_inclusive(self): + """expert_bias stays synced across gtp_remat peers only with the gtp-inclusive group.""" + if torch.cuda.device_count() < 4: + pytest.skip("Requires at least 4 CUDA devices") + _run_distributed(_worker_expert_bias_gtp_inclusive, 4) diff --git a/tests/unit_tests/generalized_tensor_parallel/test_tp_gtp.py b/tests/unit_tests/generalized_tensor_parallel/test_tp_gtp.py new file mode 100644 index 00000000000..961ab070556 --- /dev/null +++ b/tests/unit_tests/generalized_tensor_parallel/test_tp_gtp.py @@ -0,0 +1,397 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Unit tests for combined Tensor Parallelism + Generalized Tensor Parallelism (TP+GTP). + +Process group layout (world_size = tp_size x gtp_remat_size): + + rank = gtp_rank x tp_size + tp_rank + + TP group: all ranks that share the same gtp_rank (size = tp_size) + GTP group: all ranks that share the same tp_rank (size = gtp_remat_size) + +Test groups +----------- +1. TestTPGTPProcessGroups - verify TP/GTP group sizes and rank assignment +2. TestTPGTPColumnParallelLinear - column-parallel Linear: fwd/bwd correctness (weight shape verified inline) +3. TestTPGTPRowParallelLinear - row-parallel Linear: fwd/bwd smoke test + numerical correctness +4. TestTPGTPLayerNormLinear - LayerNormLinear column-parallel smoke test + +Tests use (tp_size, gtp_remat_size) = (2, 2) → world_size = 4 (runs on 4-GPU machines). + +Multi-GPU tests skip automatically when ``torch.distributed.get_world_size()`` does not match +the requested combination of tp_size x gtp_remat_size. +""" + +import types + +import pytest +import torch +import torch.distributed as dist + +from megatron.core.tensor_parallel.gtp_api import HAVE_GTP + +if not HAVE_GTP: + pytest.skip("GTP requires TransformerEngine >= 2.19", allow_module_level=True) + +import transformer_engine.pytorch as te + +from megatron.core.extensions.transformer_engine import _gtp_pre_init +from megatron.core.tensor_parallel.generalized_tensor_parallelism import ( + GTP_CONFIG, + GTPShardedParam, + update_gtp_config, + wrap_module_params_gtp, +) +from tests.unit_tests.generalized_tensor_parallel.gtp_test_utils import ( + _make_gtp_linear, + _requires_multi_gpu, + _run_distributed, + _torchrun_dist_init, + reset_fp8_state, + reset_gtp_globals, +) + + +def _build_groups(rank: int, world_size: int, tp_size: int, gtp_remat_size: int): + """Create TP and GTP process groups for a 2D parallelism grid. + + Layout: rank = gtp_rank x tp_size + tp_rank + TP group: contiguous block [gtp_rank*tp_size, (gtp_rank+1)*tp_size) + GTP group: strided set {tp_rank, tp_rank+tp_size, tp_rank+2*tp_size, ...} + + Every rank must call new_group for ALL groups (PyTorch distributed requirement). + + Returns: + tp_group: this rank's TP process group + gtp_remat_group: this rank's GTP process group + tp_rank: this rank's index within its TP group + gtp_rank: this rank's index within its GTP group + """ + assert tp_size * gtp_remat_size == world_size + tp_rank = rank % tp_size + gtp_rank = rank // tp_size + + tp_group = None + for er in range(gtp_remat_size): + ranks = list(range(er * tp_size, (er + 1) * tp_size)) + grp = dist.new_group(ranks) + if er == gtp_rank: + tp_group = grp + + gtp_remat_group = None + for tr in range(tp_size): + ranks = list(range(tr, world_size, tp_size)) + grp = dist.new_group(ranks) + if tr == tp_rank: + gtp_remat_group = grp + + return tp_group, gtp_remat_group, tp_rank, gtp_rank + + +# --------------------------------------------------------------------------- +# 1. TestTPGTPProcessGroups - group sizes and rank membership +# --------------------------------------------------------------------------- + + +def _worker_groups(rank, world_size, port, tp_size, gtp_remat_size): + tp_group, gtp_remat_group, tp_rank, gtp_rank = _build_groups( + rank, world_size, tp_size, gtp_remat_size + ) + + assert tp_group.size() == tp_size, f"rank {rank}: TP group size {tp_group.size()} != {tp_size}" + assert ( + gtp_remat_group.size() == gtp_remat_size + ), f"rank {rank}: GTP group size {gtp_remat_group.size()} != {gtp_remat_size}" + assert ( + dist.get_rank(tp_group) == tp_rank + ), f"rank {rank}: TP rank {dist.get_rank(tp_group)} != expected {tp_rank}" + assert ( + dist.get_rank(gtp_remat_group) == gtp_rank + ), f"rank {rank}: GTP rank {dist.get_rank(gtp_remat_group)} != expected {gtp_rank}" + + +class TestTPGTPProcessGroups: + @pytest.mark.parametrize("tp_size,gtp_remat_size", [(2, 2)]) + def test_group_sizes_and_ranks(self, tp_size, gtp_remat_size): + world_size = tp_size * gtp_remat_size + _requires_multi_gpu(world_size) + _run_distributed(_worker_groups, world_size, tp_size, gtp_remat_size) + + +# --------------------------------------------------------------------------- +# 2. TestTPGTPColumnParallelLinear +# --------------------------------------------------------------------------- + + +def _worker_column_correctness(rank, world_size, port, tp_size, gtp_remat_size): + """Column-parallel output must equal inp @ (GTP-gathered TP-local weight)^T.""" + torch.manual_seed(0) + tp_group, gtp_remat_group, tp_rank, gtp_rank = _build_groups( + rank, world_size, tp_size, gtp_remat_size + ) + + batch, in_f = 16, 64 + out_f = tp_size * gtp_remat_size * 32 # per-rank shard = 32 rows + dtype = torch.bfloat16 + + layer = _make_gtp_linear( + in_f, out_f, gtp_remat_group, dtype, parallel_mode="column", tp_group=tp_group + ) + + # All-gather GTP_remat shards → TP-local full weight [out_f/tp_size, in_f] + shard = layer.weight.data.clone() + all_gtp_shards = [torch.zeros_like(shard) for _ in range(gtp_remat_size)] + dist.all_gather(all_gtp_shards, shard, group=gtp_remat_group) + tp_local_weight = torch.cat(all_gtp_shards, dim=0).float() # strip padding + tp_local_weight = tp_local_weight[: out_f // tp_size] + + # Same full input on all ranks (column-parallel: each rank processes full input) + inp = torch.randn(batch, in_f, dtype=dtype, device="cuda") + dist.broadcast(inp, src=0) + inp_te = inp.clone().requires_grad_(True) + + # TE forward: GTP_remat all-gathers weight internally; no TP comm in column-parallel fwd + out = layer(inp_te, is_first_microbatch=True) + assert out.shape == ( + batch, + out_f // tp_size, + ), f"rank {rank}: output shape {out.shape} != ({batch}, {out_f // tp_size})" + + # Reference: this TP rank's output = inp @ tp_local_weight^T + ref = inp.float() @ tp_local_weight.T + ref = ref.to(dtype) + assert torch.allclose( + out.float(), ref.float(), atol=1e-2, rtol=1e-2 + ), f"rank {rank}: output mismatch, max_diff={(out.float() - ref.float()).abs().max():.4f}" + + # Backward: dX is all-reduced across TP group internally by TE + grad = torch.randn_like(out) + dist.broadcast(grad, src=0) + # wgrad RS path always accumulates into main_grad; allocate before backward. + layer.weight.main_grad = torch.zeros(layer.weight.shape, dtype=dtype, device="cuda") + out.backward(grad) + assert inp_te.grad is not None and inp_te.grad.shape == inp.shape + assert torch.isfinite(inp_te.grad).all(), f"rank {rank}: non-finite dX" + + +class TestTPGTPColumnParallelLinear: + @pytest.mark.parametrize("tp_size,gtp_remat_size", [(2, 2)]) + def test_forward_backward_correctness(self, tp_size, gtp_remat_size): + world_size = tp_size * gtp_remat_size + _requires_multi_gpu(world_size) + _run_distributed(_worker_column_correctness, world_size, tp_size, gtp_remat_size) + + +# --------------------------------------------------------------------------- +# 3. TestTPGTPRowParallelLinear +# --------------------------------------------------------------------------- + + +def _worker_row_forward_backward(rank, world_size, port, tp_size, gtp_remat_size): + """Row-parallel: weight shape verified; output is all-reduced [batch, out_f]; backward produces finite dX.""" + torch.manual_seed(0) + tp_group, gtp_remat_group, tp_rank, _ = _build_groups(rank, world_size, tp_size, gtp_remat_size) + + batch = 16 + in_f = tp_size * 64 # full in_features + out_f = gtp_remat_size * 64 # full out_features + dtype = torch.bfloat16 + + layer = _make_gtp_linear( + in_f, out_f, gtp_remat_group, dtype, parallel_mode="row", tp_group=tp_group + ) + + expected_shape = (out_f // gtp_remat_size, in_f // tp_size) + assert isinstance( + layer.weight, GTPShardedParam + ), f"rank {rank}: weight should be GTPShardedParam" + assert ( + layer.weight.shape == expected_shape + ), f"rank {rank}: expected {expected_shape}, got {layer.weight.shape}" + + # Row-parallel: each TP rank takes the corresponding slice of in_f + full_inp = torch.randn(batch, in_f, dtype=dtype, device="cuda") + dist.broadcast(full_inp, src=0) + local_in_f = in_f // tp_size + inp = full_inp[:, tp_rank * local_in_f : (tp_rank + 1) * local_in_f] + inp = inp.clone().requires_grad_(True) + + # TE forward: GTP_remat all-gathers weight, row-parallel all-reduces output across TP + out = layer(inp, is_first_microbatch=True) + assert out.shape == ( + batch, + out_f, + ), f"rank {rank}: output shape {out.shape} != ({batch}, {out_f})" + assert torch.isfinite(out).all(), f"rank {rank}: non-finite output" + + # wgrad RS path always accumulates into main_grad; allocate before backward. + layer.weight.main_grad = torch.zeros(layer.weight.shape, dtype=dtype, device="cuda") + out.sum().backward() + assert inp.grad is not None and inp.grad.shape == inp.shape + assert torch.isfinite(inp.grad).all(), f"rank {rank}: non-finite dX" + + +def _worker_row_correctness(rank, world_size, port, tp_size, gtp_remat_size): + """Row-parallel all-reduced output must equal inp_full @ full_weight^T.""" + torch.manual_seed(0) + tp_group, gtp_remat_group, tp_rank, _ = _build_groups(rank, world_size, tp_size, gtp_remat_size) + + batch = 16 + in_f = tp_size * 64 + out_f = gtp_remat_size * 64 + dtype = torch.bfloat16 + + layer = _make_gtp_linear( + in_f, out_f, gtp_remat_group, dtype, parallel_mode="row", tp_group=tp_group + ) + + # Reconstruct full weight: all-gather GTP_remat shards → TP-local, then all-gather TP shards + shard = layer.weight.data.clone() + all_gtp_shards = [torch.zeros_like(shard) for _ in range(gtp_remat_size)] + dist.all_gather(all_gtp_shards, shard, group=gtp_remat_group) + tp_local_weight = torch.cat(all_gtp_shards, dim=0).float() # [out_f, in_f/tp_size] + + all_tp_weights = [torch.zeros_like(tp_local_weight) for _ in range(tp_size)] + dist.all_gather(all_tp_weights, tp_local_weight, group=tp_group) + full_weight = torch.cat(all_tp_weights, dim=1).float() # [out_f, in_f] + + # Full input (same on all ranks; we slice below to simulate row-parallel) + full_inp = torch.randn(batch, in_f, dtype=dtype, device="cuda") + dist.broadcast(full_inp, src=0) + local_in_f = in_f // tp_size + inp = full_inp[:, tp_rank * local_in_f : (tp_rank + 1) * local_in_f].clone() + inp.requires_grad_(True) + + out = layer(inp, is_first_microbatch=True) + + # Reference: full input @ full weight^T — all ranks should see the same output + ref = full_inp.float() @ full_weight.T + ref = ref.to(dtype) + assert torch.allclose( + out.float(), ref.float(), atol=2e-2, rtol=1e-2 + ), f"rank {rank}: output mismatch, max_diff={(out.float() - ref.float()).abs().max():.4f}" + + +class TestTPGTPRowParallelLinear: + @pytest.mark.parametrize("tp_size,gtp_remat_size", [(2, 2)]) + def test_forward_backward(self, tp_size, gtp_remat_size): + world_size = tp_size * gtp_remat_size + _requires_multi_gpu(world_size) + _run_distributed(_worker_row_forward_backward, world_size, tp_size, gtp_remat_size) + + @pytest.mark.parametrize("tp_size,gtp_remat_size", [(2, 2)]) + def test_forward_correctness(self, tp_size, gtp_remat_size): + world_size = tp_size * gtp_remat_size + _requires_multi_gpu(world_size) + _run_distributed(_worker_row_correctness, world_size, tp_size, gtp_remat_size) + + +# --------------------------------------------------------------------------- +# 4. TestTPGTPLayerNormLinear - column-parallel smoke test +# --------------------------------------------------------------------------- + + +def _worker_layernorm_linear(rank, world_size, port, tp_size, gtp_remat_size): + torch.manual_seed(0) + tp_group, gtp_remat_group, _, _ = _build_groups(rank, world_size, tp_size, gtp_remat_size) + + seq, batch = 4, 2 + in_f = 64 + out_f = tp_size * gtp_remat_size * 32 + dtype = torch.bfloat16 + + layer = te.LayerNormLinear( + in_features=in_f, + out_features=out_f, + bias=False, + params_dtype=dtype, + parallel_mode="column", + device="cuda", + tp_group=tp_group, + ) + # TE has no GTP construction hook; gtp_remat_size (the forward-gather gate) is stamped + # post-init and the TP-sharded weight is sliced into a GTPShardedParam on the Megatron side. + layer.gtp_remat_size = gtp_remat_group.size() + wrap_module_params_gtp(layer, layer.weight_names, gtp_remat_group) + assert isinstance( + layer.weight, GTPShardedParam + ), f"rank {rank}: LayerNormLinear.weight should be GTPShardedParam" + expected_rows = out_f // (tp_size * gtp_remat_size) + assert layer.weight.shape == ( + expected_rows, + in_f, + ), f"rank {rank}: unexpected weight shape {layer.weight.shape}" + + inp = torch.randn(seq, batch, in_f, dtype=dtype, device="cuda", requires_grad=True) + dist.broadcast(inp, src=0) + + out = layer(inp, is_first_microbatch=True) + assert out.shape == (seq, batch, out_f // tp_size), f"rank {rank}: output shape {out.shape}" + assert torch.isfinite(out).all(), f"rank {rank}: non-finite output" + + # wgrad RS path always accumulates into main_grad; allocate before backward. + layer.weight.main_grad = torch.zeros(layer.weight.shape, dtype=dtype, device="cuda") + out.sum().backward() + assert inp.grad is not None and inp.grad.shape == inp.shape + assert torch.isfinite(inp.grad).all(), f"rank {rank}: non-finite dX" + + +class TestTPGTPLayerNormLinear: + @pytest.mark.parametrize("tp_size,gtp_remat_size", [(2, 2)]) + def test_forward_backward(self, tp_size, gtp_remat_size): + world_size = tp_size * gtp_remat_size + _requires_multi_gpu(world_size) + _run_distributed(_worker_layernorm_linear, world_size, tp_size, gtp_remat_size) + + +# --------------------------------------------------------------------------- +# 5. TestTPGTPPaddingAlignment - GTP pre-shard must pad the *per-TP* slice so the weight +# stays MXFP8-aligned AFTER TE's tp-split (padding the full out_features would let the +# tp-split de-align it). +# --------------------------------------------------------------------------- + + +def _worker_pre_init_tp_padding(rank, world_size, port, tp_size, gtp_remat_size): + """`_gtp_pre_init` pads ``output_size // out_split_size`` (the per-TP slice), so the shard + TE hands each rank after its own tp-split is a multiple of ``pad_for_alignment`` (32 for MXFP8). + + tp2 x gtp2, pad=32, out_features=192: + correct: per-TP slice 96 -> pad 128 -> /gtp2 = 64 -> TE tp-split /2 = 64 (32-aligned). + pre-fix: pad full 192 -> /gtp2 = 96 -> TE tp-split /2 = 48 (NOT 32-aligned). + """ + tp_group, gtp_remat_group, tp_rank, gtp_rank = _build_groups( + rank, world_size, tp_size, gtp_remat_size + ) + out_features = 192 + orig_pad = GTP_CONFIG.pad_for_alignment + update_gtp_config(pad_for_alignment=32) + try: + extra_kwargs = {} + shard_out, gtp_ctx = _gtp_pre_init( + types.SimpleNamespace(), + out_features, + gtp_remat_group, + extra_kwargs, + out_split_size=tp_size, + ) + # shard_out is what TE receives as out_features; TE splits it again across tp_size. + per_gpu = shard_out // tp_size + assert per_gpu % 32 == 0, ( + f"rank {rank}: per-GPU shard {per_gpu} not 32-aligned -- TE tp-split de-aligned the " + "GTP pad (pad the per-TP slice, not the full out_features)" + ) + assert per_gpu == 64, f"rank {rank}: expected per-GPU 64, got {per_gpu}" + # pad_length is measured on the per-TP slice: 128 - 96 = 32. + gtp_remat_group_ctx, pad_length, logical = gtp_ctx + assert pad_length == 32, f"rank {rank}: expected pad_length 32, got {pad_length}" + assert gtp_remat_group_ctx.size() == gtp_remat_size and logical == out_features + finally: + update_gtp_config(pad_for_alignment=orig_pad) + + +class TestTPGTPPaddingAlignment: + @pytest.mark.parametrize("tp_size,gtp_remat_size", [(2, 2)]) + def test_pre_init_pads_per_tp_slice(self, tp_size, gtp_remat_size): + world_size = tp_size * gtp_remat_size + _requires_multi_gpu(world_size) + _run_distributed(_worker_pre_init_tp_padding, world_size, tp_size, gtp_remat_size) diff --git a/tests/unit_tests/inference/contexts/attention_metadata/test_mamba_metadata.py b/tests/unit_tests/inference/contexts/attention_metadata/test_mamba_metadata.py index a7f579051ac..99bb046d97d 100644 --- a/tests/unit_tests/inference/contexts/attention_metadata/test_mamba_metadata.py +++ b/tests/unit_tests/inference/contexts/attention_metadata/test_mamba_metadata.py @@ -14,7 +14,14 @@ def metadata_context(self): """Fixture to initialize MambaMetadata with standard constraints.""" max_requests = 16 max_tokens = 2048 - metadata = MambaMetadata(max_requests=max_requests, max_tokens=max_tokens) + # Per-step intermediate-state cap (token budget / block_size + margin); + # value is irrelevant to these update() tests, which don't extract state. + max_intermediate_count = 17 + metadata = MambaMetadata( + max_requests=max_requests, + max_tokens=max_tokens, + max_intermediate_count=max_intermediate_count, + ) # Manually allocate some slots to simulate a running state. # We assume request_id i maps to mamba_slot i for simplicity in assertions. diff --git a/tests/unit_tests/inference/contexts/test_dynamic_context.py b/tests/unit_tests/inference/contexts/test_dynamic_context.py index 0008ba043b7..21ca613512c 100644 --- a/tests/unit_tests/inference/contexts/test_dynamic_context.py +++ b/tests/unit_tests/inference/contexts/test_dynamic_context.py @@ -2,6 +2,7 @@ import contextlib import math +from types import SimpleNamespace from unittest import mock import pytest @@ -37,6 +38,22 @@ def rounder_override(n): DynamicInferenceContext.REQUEST_ROUNDER = original_request_rounder +@pytest.mark.parametrize( + "using_cuda_graph, num_prefill_requests, padded_prefill_requests, expected", + [(False, 0, 1, True), (False, 1, 0, False), (True, 1, 0, True), (True, 0, 1, False)], +) +def test_is_decode_only_uses_current_execution_snapshot( + using_cuda_graph, num_prefill_requests, padded_prefill_requests, expected +): + """Decode-only classification follows the eager or CUDA graph execution state.""" + context = DynamicInferenceContext.__new__(DynamicInferenceContext) + context._using_cuda_graph_this_step = using_cuda_graph + context.num_prefill_requests = num_prefill_requests + context.padded_batch_dimensions = mock.Mock(prefill_req_count=padded_prefill_requests) + + assert context.is_decode_only() is expected + + class TestDynamicContext: @classmethod @@ -298,6 +315,147 @@ def test_current_input_and_position_ids_view_cache(self): assert torch.equal(refreshed_input_ids.squeeze(0), new_input_ids) assert torch.equal(refreshed_pos_ids.squeeze(0), new_pos_ids) + @pytest.mark.internal + @rounder_override(64) + @pytest.mark.parametrize( + "transfer_bookkeeping,record_done_event,expected_event", + [(False, False, None), (True, False, None), (True, True, "bookkeeping")], + ) + def test_initialize_attention_state_bookkeeping_transfer_event( + self, transfer_bookkeeping, record_done_event, expected_event + ): + dynamic_context = self._get_dynamic_context( + params_dtype=torch.float32, + num_layers=2, + kv_channels=64, + num_attention_heads=8, + max_sequence_length=128, + buffer_size_gb=0.1, + block_size_tokens=128, + max_tokens=None, + ) + dynamic_context.transfer_bookkeeping_to_gpu = mock.Mock( + side_effect=lambda *, record_done_event=False: ( + "bookkeeping" if record_done_event else None + ) + ) + + done_event = dynamic_context.initialize_attention_state( + transfer_bookkeeping_to_gpu=transfer_bookkeeping, + record_bookkeeping_done_event=record_done_event, + ) + + if transfer_bookkeeping: + dynamic_context.transfer_bookkeeping_to_gpu.assert_called_once_with( + record_done_event=record_done_event + ) + else: + dynamic_context.transfer_bookkeeping_to_gpu.assert_not_called() + assert done_event == expected_event + + @pytest.mark.internal + @rounder_override(64) + def test_transfer_bookkeeping_to_gpu_can_skip_input_token_ids(self): + dynamic_context = self._get_dynamic_context( + params_dtype=torch.float32, + num_layers=2, + kv_channels=64, + num_attention_heads=8, + max_sequence_length=128, + buffer_size_gb=0.1, + block_size_tokens=128, + max_tokens=None, + ) + + num_tokens = 4 + dynamic_context.total_request_count = 2 + dynamic_context.paused_request_count = 0 + dynamic_context.padded_active_request_count = 2 + dynamic_context.token_to_input_ids[:num_tokens] = torch.tensor( + [11, 12, 13, 14], dtype=torch.int64 + ) + dynamic_context.token_to_pos_ids[:num_tokens] = torch.tensor( + [21, 22, 23, 24], dtype=torch.int64 + ) + existing_gpu_tokens = torch.tensor( + [91, 92, 93, 94], + dtype=torch.int64, + device=dynamic_context.gpu_view.token_to_input_ids.device, + ) + dynamic_context.gpu_view.token_to_input_ids[:num_tokens] = existing_gpu_tokens + + done_event = dynamic_context.transfer_bookkeeping_to_gpu( + skip_token_input_ids=True, record_done_event=True + ) + done_event.synchronize() + + assert torch.equal( + dynamic_context.gpu_view.token_to_input_ids[:num_tokens], existing_gpu_tokens + ) + assert torch.equal( + dynamic_context.gpu_view.token_to_pos_ids[:num_tokens].cpu(), + torch.tensor([21, 22, 23, 24], dtype=torch.int64), + ) + + dynamic_context.token_to_input_ids[:num_tokens] = torch.tensor( + [31, 32, 33, 34], dtype=torch.int64 + ) + dynamic_context.transfer_bookkeeping_to_gpu() + + assert torch.equal( + dynamic_context.gpu_view.token_to_input_ids[:num_tokens].cpu(), + torch.tensor([31, 32, 33, 34], dtype=torch.int64), + ) + + @pytest.mark.internal + @rounder_override(8) + @pytest.mark.parametrize("num_speculative_tokens", [0, 2]) + def test_copy_async_sched_sample_to_forward_populates_active_and_clears_padding( + self, num_speculative_tokens + ): + ctx = self._get_dynamic_context( + params_dtype=torch.float32, + num_layers=2, + kv_channels=8, + num_attention_heads=2, + max_sequence_length=32, + buffer_size_gb=0.01, + block_size_tokens=4, + max_tokens=32, + max_requests=8, + num_speculative_tokens=num_speculative_tokens, + ) + + ctx.total_request_count = 3 + ctx.paused_request_count = 0 + ctx.num_prefill_requests = 0 + token_count = 3 * (num_speculative_tokens + 1) + ctx.active_token_count = token_count + ctx.padded_active_token_count = 12 + device = ctx.gpu_view.token_to_input_ids.device + ctx.gpu_view.token_to_input_ids[:12] = torch.full( + (12,), 777, dtype=torch.int64, device=device + ) + sampled_tokens_cuda = torch.tensor([90, 91, 92], dtype=torch.int64, device=device) + sampled_mtp_tokens_cuda = ( + torch.tensor([[100, 101, 102], [110, 111, 112]], device=device) + if num_speculative_tokens > 0 + else None + ) + + ctx.copy_async_sched_sample_to_forward(sampled_tokens_cuda, sampled_mtp_tokens_cuda) + + expected_tokens = ( + sampled_tokens_cuda + if sampled_mtp_tokens_cuda is None + else torch.tensor([90, 100, 110, 91, 101, 111, 92, 102, 112], device=device) + ) + assert torch.equal(ctx.gpu_view.token_to_input_ids[:token_count], expected_tokens) + assert torch.equal( + ctx.gpu_view.token_to_input_ids[token_count:12].cpu(), + torch.zeros(12 - token_count, dtype=torch.int64), + ) + @pytest.mark.internal @rounder_override(64) @pytest.mark.parametrize("is_hybrid_model", [False, True]) @@ -318,6 +476,9 @@ def test_reset(self, is_hybrid_model: bool): # Initialize all variables dynamic_context.total_request_count = 10 dynamic_context.active_token_count = 10 + dynamic_context.step_count = 4 + dynamic_context.prefix_cache_lru_clock = 5 + dynamic_context.lifetime_prefill_token_count = 6 dynamic_context.async_sched_step_count = 6 dynamic_context.async_sched_compaction_step_count = 7 dynamic_context.paused_request_count = 5 @@ -348,6 +509,9 @@ def test_reset(self, is_hybrid_model: bool): # Assert all variables are reset to zero or their default values assert dynamic_context.total_request_count == 0 assert dynamic_context.active_token_count == 0 + assert dynamic_context.step_count == 0 + assert dynamic_context.prefix_cache_lru_clock == 0 + assert dynamic_context.lifetime_prefill_token_count == 0 assert dynamic_context.async_sched_step_count == 0 assert dynamic_context.async_sched_compaction_step_count == 0 assert dynamic_context.paused_request_count == 0 @@ -854,7 +1018,7 @@ def test_update_request(self, is_hybrid_model: bool): ) ) - def _get_async_sched_context(self): + def _get_async_sched_context(self, num_speculative_tokens=0, is_hybrid_model=False): return self._get_dynamic_context( params_dtype=torch.float32, num_layers=2, @@ -865,6 +1029,9 @@ def _get_async_sched_context(self): block_size_tokens=4, max_tokens=32, max_requests=8, + num_speculative_tokens=num_speculative_tokens, + is_hybrid_model=is_hybrid_model, + layer_type_list=[Symbols.MAMBA, Symbols.ATTENTION], ) @staticmethod @@ -884,6 +1051,7 @@ def _setup_async_sched_decode_rows( active_slice = slice(0, active_request_count) ctx.request_ids[active_slice] = torch.tensor(request_ids, dtype=torch.int32) + ctx.request_in_prefill_status_tensor[active_slice] = 0 ctx.request_query_lengths[active_slice] = 1 ctx.request_output_lengths[active_slice] = 16 ctx.request_kv_length_offsets[active_slice] = torch.tensor(kv_offsets, dtype=torch.int32) @@ -907,48 +1075,53 @@ def _setup_async_sched_decode_rows( ctx.token_to_local_position_within_kv_block[active_slice] = ( ctx.token_to_pos_ids[active_slice] % ctx.block_size_tokens ) + if ctx.is_hybrid_model: + mamba_slots = ctx.mamba_metadata.batch_allocate_slots(active_request_count) + assert mamba_slots is not None + ctx.mamba_metadata.request_to_mamba_state_idx[active_slice] = mamba_slots @pytest.mark.internal @rounder_override(8) @pytest.mark.parametrize( - "new_tokens, kv_offsets, last_offsets, expected_kv_offsets, expected_last_offsets", + "active_request_count, kv_offsets, last_offsets, expected_kv_offsets, expected_last_offsets", [ - ([], [], [], [], []), - ([90, 91], [3, 5], [1, 2], [4, 6], [2, 3]), - ([90, 91], [3, 5], [3, 1], [4, 6], [0, 2]), + (0, [], [], [], []), + (2, [3, 5], [1, 2], [4, 6], [2, 3]), + (2, [3, 5], [3, 1], [4, 6], [0, 2]), ], ) def test_async_sched_prepare_requests_success( - self, new_tokens, kv_offsets, last_offsets, expected_kv_offsets, expected_last_offsets + self, + active_request_count, + kv_offsets, + last_offsets, + expected_kv_offsets, + expected_last_offsets, ): """Async scheduling prepare advances active decode rows without lifecycle changes.""" ctx = self._get_async_sched_context() self._setup_async_sched_decode_rows( ctx, - active_request_count=len(new_tokens), + active_request_count=active_request_count, kv_offsets=kv_offsets, last_block_offsets=last_offsets, ) - tokens = torch.tensor(new_tokens, dtype=torch.int64) - if new_tokens and torch.cuda.is_available(): - tokens = tokens.cuda() + original_tokens = ctx.token_to_input_ids[:active_request_count].clone() - ctx.prepare_requests(tokens) + ctx.prepare_requests() - assert ctx.active_token_count == len(new_tokens) + assert ctx.active_token_count == active_request_count assert torch.equal( - ctx.request_kv_length_offsets[: len(new_tokens)], + ctx.request_kv_length_offsets[:active_request_count], torch.tensor(expected_kv_offsets, dtype=torch.int32), ) assert torch.equal( - ctx.request_last_kv_block_offset[: len(new_tokens)], + ctx.request_last_kv_block_offset[:active_request_count], torch.tensor(expected_last_offsets, dtype=torch.int32), ) + assert torch.equal(ctx.token_to_input_ids[:active_request_count], original_tokens) assert torch.equal( - ctx.token_to_input_ids[: len(new_tokens)], torch.tensor(new_tokens, dtype=torch.long) - ) - assert torch.equal( - ctx.token_to_pos_ids[: len(new_tokens)], + ctx.token_to_pos_ids[:active_request_count], torch.tensor(expected_kv_offsets, dtype=torch.long), ) if last_offsets and last_offsets[0] == ctx.block_size_tokens - 1: @@ -958,17 +1131,113 @@ def test_async_sched_prepare_requests_success( @pytest.mark.internal @rounder_override(8) @pytest.mark.parametrize( - "setup, new_tokens, expected_message", + "num_speculative_tokens, last_block_offsets, active_avail, expected", + [ + (0, [0, 1], 0, True), + (0, [3, 1], 0, False), + (0, [3, 1], 1, True), + (2, [1, 2], 1, False), + (2, [1, 2], 2, True), + ], + ) + def test_async_sched_can_prepare_requests_exact_block_demand( + self, num_speculative_tokens, last_block_offsets, active_avail, expected + ): + """Overlap capacity counts only requests crossing a block boundary.""" + ctx = self._get_async_sched_context(num_speculative_tokens=num_speculative_tokens) + self._setup_async_sched_decode_rows( + ctx, active_request_count=len(last_block_offsets), last_block_offsets=last_block_offsets + ) + ctx.kv_block_allocator.get_active_avail = mock.Mock(return_value=active_avail) + + assert ctx.can_prepare_requests() is expected + + @pytest.mark.internal + @rounder_override(8) + @pytest.mark.parametrize("state", ["prefill", "paused"]) + def test_async_sched_cannot_prepare_requests_with_lifecycle_state(self, state): + """Overlap preparation rejects state requiring lifecycle bookkeeping.""" + ctx = self._get_async_sched_context() + self._setup_async_sched_decode_rows(ctx, active_request_count=2) + if state == "prefill": + ctx.num_prefill_requests = 1 + else: + ctx.paused_request_count = 1 + + assert not ctx.can_prepare_requests() + + @pytest.mark.internal + @rounder_override(8) + def test_async_sched_prepare_capacity_recovers_after_pause_resume(self): + """No-overlap bookkeeping restores overlap eligibility after resuming a request.""" + ctx = self._get_async_sched_context() + self._setup_async_sched_decode_rows( + ctx, active_request_count=2, last_block_offsets=[ctx.block_size_tokens - 1, 0] + ) + ctx.kv_block_allocator.active_count = ctx.kv_block_allocator.get_active_used() + ctx.kv_block_allocator.total_avail = 0 + ctx.kv_block_allocator.paused_count = 100 + + assert not ctx.can_prepare_requests() + + ctx.update_requests( + active_requests_mask=torch.tensor([1, 1]), new_tokens=torch.tensor([90, 91]) + ) + + assert ctx.paused_request_count == 1 + assert not ctx.can_prepare_requests() + + ctx.kv_block_allocator.total_avail = 1 + ctx.update_requests(active_requests_mask=torch.tensor([0]), new_tokens=torch.tensor([92])) + + assert ctx.paused_request_count == 0 + assert ctx.can_prepare_requests() + + @pytest.mark.internal + @rounder_override(8) + def test_async_sched_commit_sampled_tokens(self): + """Async scheduling commits sampled CPU tokens after prepare.""" + ctx = self._get_async_sched_context() + self._setup_async_sched_decode_rows(ctx, active_request_count=2, kv_offsets=[3, 5]) + original_tokens = ctx.token_to_input_ids[:2].clone() + + ctx.prepare_requests() + + assert torch.equal(ctx.token_to_input_ids[:2], original_tokens) + assert torch.equal( + ctx.request_kv_length_offsets[:2], torch.tensor([4, 6], dtype=torch.int32) + ) + + sampled_tokens_cpu = torch.tensor([90, 91], dtype=torch.int64) + if torch.cuda.is_available(): + with pytest.raises(AssertionError, match="must be on the CPU"): + ctx.commit_sampled_tokens(sampled_tokens_cpu.cuda()) + ctx.active_token_count = 5 + ctx.commit_sampled_tokens(sampled_tokens_cpu) + + assert ctx.active_token_count == 2 + assert torch.equal(ctx.token_to_input_ids[:2], sampled_tokens_cpu) + + with pytest.raises(RuntimeError, match="Expected 2 new tokens"): + ctx.commit_sampled_tokens(torch.tensor([90], dtype=torch.int64)) + + ctx.total_request_count = 0 + ctx.active_token_count = 2 + ctx.commit_sampled_tokens(torch.empty(0, dtype=torch.int64)) + assert ctx.active_token_count == 0 + + @pytest.mark.internal + @rounder_override(8) + @pytest.mark.parametrize( + "setup, expected_message", [ - (lambda ctx: setattr(ctx, "num_speculative_tokens", 1), [90, 91], "speculative"), - (lambda ctx: setattr(ctx, "num_prefill_requests", 1), [90, 91], "decode-only"), - (lambda ctx: setattr(ctx, "paused_request_count", 1), [90, 91], "paused"), - (lambda ctx: None, [90], "Expected 2 new tokens"), - (lambda ctx: None, [90, 91], "pause requests"), - (lambda ctx: None, [90, 91], "evict requests"), + (lambda ctx: setattr(ctx, "num_prefill_requests", 1), "decode-only"), + (lambda ctx: setattr(ctx, "paused_request_count", 1), "paused"), + (lambda ctx: None, "pause requests"), + (lambda ctx: None, "evict requests"), ], ) - def test_async_sched_prepare_requests_errors(self, setup, new_tokens, expected_message): + def test_async_sched_prepare_requests_errors(self, setup, expected_message): """Async scheduling prepare raises instead of performing lifecycle operations.""" ctx = self._get_async_sched_context() self._setup_async_sched_decode_rows( @@ -983,19 +1252,30 @@ def test_async_sched_prepare_requests_errors(self, setup, new_tokens, expected_m setup(ctx) with pytest.raises(RuntimeError, match=expected_message): - ctx.prepare_requests(torch.tensor(new_tokens, dtype=torch.int64)) + ctx.prepare_requests() @pytest.mark.internal @rounder_override(8) + @pytest.mark.parametrize("is_hybrid_model", [False, True]) @pytest.mark.parametrize( - "mask, expected_finished_ids, expected_request_ids", - [([1, 1, 1], [], [10, 11, 12]), ([1, 0, 1], [11], [10, 12]), ([0, 0, 0], [10, 11, 12], [])], + "mask, expected_finished_ids, expected_request_ids, expected_survivor_idxs", + [ + ([1, 1, 1], [], [10, 11, 12], [0, 1, 2]), + ([1, 0, 1], [11], [10, 12], [0, 2]), + ([0, 1, 1], [10], [12, 11], [2, 1]), + ([0, 0, 0], [10, 11, 12], [], []), + ], ) def test_async_sched_resolve_requests_success( - self, mask, expected_finished_ids, expected_request_ids + self, + mask, + expected_finished_ids, + expected_request_ids, + expected_survivor_idxs, + is_hybrid_model, ): """Async scheduling resolve compacts survivors and releases finished rows.""" - ctx = self._get_async_sched_context() + ctx = self._get_async_sched_context(is_hybrid_model=is_hybrid_model) self._setup_async_sched_decode_rows( ctx, active_request_count=len(mask), @@ -1003,35 +1283,131 @@ def test_async_sched_resolve_requests_success( kv_offsets=[4, 5, 6], last_block_offsets=[0, 1, 2], ) + original_mamba_slots = ( + ctx.mamba_metadata.request_to_mamba_state_idx[: len(mask)].clone() + if is_hybrid_model + else None + ) + mamba_state_bank_ptrs = ( + (ctx.mamba_conv_states.data_ptr(), ctx.mamba_ssm_states.data_ptr()) + if is_hybrid_model + else None + ) active_mask = torch.tensor(mask, dtype=torch.int32) if torch.cuda.is_available(): active_mask = active_mask.cuda() - finished_request_ids = ctx.resolve_requests(active_mask) + token_tensors = ( + ctx.token_to_input_ids, + ctx.token_to_pos_ids, + ctx.token_to_block_idx, + ctx.token_to_local_position_within_kv_block, + ctx.token_to_request_idx, + ctx.token_to_position_in_request, + ) + active_token_count = ctx.active_token_count + token_state = tuple(tensor.clone() for tensor in token_tensors) + + finished_request_ids, survivor_idxs = ctx.resolve_requests(active_mask) assert torch.equal( finished_request_ids, torch.tensor(expected_finished_ids, dtype=torch.int32) ) + assert torch.equal(survivor_idxs, torch.tensor(expected_survivor_idxs)) assert ctx.total_request_count == len(expected_request_ids) - assert ctx.active_token_count == len(expected_request_ids) + assert ctx.active_token_count == active_token_count assert torch.equal( ctx.request_ids[: len(expected_request_ids)], torch.tensor(expected_request_ids, dtype=torch.int32), ) - assert torch.equal( - ctx.token_to_request_idx[: len(expected_request_ids)], - torch.arange(len(expected_request_ids), dtype=torch.int32), - ) + for tensor, expected in zip(token_tensors, token_state): + assert torch.equal(tensor, expected) if not expected_request_ids: assert torch.all(ctx.request_to_kv_block_ids == -1) + if is_hybrid_model: + expected_mamba_slots = original_mamba_slots[survivor_idxs] + assert torch.equal( + ctx.mamba_metadata.request_to_mamba_state_idx[: len(expected_request_ids)], + expected_mamba_slots, + ) + assert torch.all( + ctx.mamba_metadata.request_to_mamba_state_idx[len(expected_request_ids) : len(mask)] + == -1 + ) + assert mamba_state_bank_ptrs == ( + ctx.mamba_conv_states.data_ptr(), + ctx.mamba_ssm_states.data_ptr(), + ) + + @pytest.mark.internal + @rounder_override(8) + def test_async_sched_mtp_prepare_commit_and_resolve(self): + """MTP survivor tokens are committed after request resolution.""" + ctx = self._get_async_sched_context(num_speculative_tokens=2) + self._setup_async_sched_decode_rows( + ctx, + active_request_count=2, + request_ids=[10, 11], + kv_offsets=[3, 5], + last_block_offsets=[1, 3], + ) + + ctx.prepare_requests() + prepared_input_ids = ctx.token_to_input_ids.clone() + + assert ctx.active_token_count == 6 + assert torch.equal(ctx.token_to_pos_ids[:6], torch.tensor([4, 5, 6, 6, 7, 8])) + + finished_request_ids, survivor_idxs = ctx.resolve_requests(torch.tensor([0, 1])) + + assert finished_request_ids.tolist() == [10] + assert survivor_idxs.tolist() == [1] + assert ctx.request_ids[0] == 11 + assert ctx.active_token_count == 6 + assert torch.equal(ctx.token_to_input_ids, prepared_input_ids) + + sampled_tokens = torch.tensor([100, 200]) + sampled_mtp_tokens = torch.tensor([[101, 201], [102, 202]]) + ctx.commit_sampled_tokens( + sampled_tokens[survivor_idxs], sampled_mtp_tokens[:, survivor_idxs] + ) + + assert ctx.active_token_count == 3 + assert torch.equal(ctx.token_to_input_ids[:3], torch.tensor([200, 201, 202])) + + @pytest.mark.internal + @rounder_override(8) + def test_async_sched_prefill_resolves_before_decode_prepare(self): + """Resolution converts prefill survivors before prepare rebuilds decode rows.""" + ctx = self._get_async_sched_context() + self._setup_async_sched_decode_rows( + ctx, + active_request_count=2, + request_ids=[10, 11], + kv_offsets=[4, 6], + last_block_offsets=[0, 2], + ) + ctx.num_prefill_requests = 1 + ctx.request_in_prefill_status_tensor[1] = 1 + ctx.request_query_lengths[1] = 4 + ctx.active_token_count = 5 + + _, survivor_idxs = ctx.resolve_requests(torch.tensor([1, 1])) + assert ctx.active_token_count == 5 + + ctx.prepare_requests() + + assert survivor_idxs.tolist() == [0, 1] + assert ctx.num_prefill_requests == 0 + assert ctx.active_token_count == 2 + assert torch.equal(ctx.request_query_lengths[:2], torch.tensor([1, 1])) + assert torch.equal(ctx.request_kv_length_offsets[:2], torch.tensor([5, 10])) @pytest.mark.internal @rounder_override(8) @pytest.mark.parametrize( "setup, mask, expected_message", [ - (lambda ctx: setattr(ctx, "num_speculative_tokens", 1), [1, 1], "speculative"), - (lambda ctx: setattr(ctx, "num_prefill_requests", 1), [1, 1], "decode-only"), (lambda ctx: setattr(ctx, "paused_request_count", 1), [1, 1], "paused"), (lambda ctx: None, [1], "Expected active mask"), ], @@ -1337,22 +1713,38 @@ def expected_log_probs(logits, active_id_and_counts): For processed mode, each active request's params are repeated across its token count, mirroring the request->row mapping in `_processed_log_probs`. + + Args: + logits (Tensor): Raw logits for the active token rows. + active_id_and_counts: Request IDs paired with their row counts. + + Returns: + Tensor: Expected raw or sampling-processed log probabilities. """ logits_2d = logits.squeeze(0).float() if logprobs_mode == "raw_logprobs": return torch.nn.functional.log_softmax(logits_2d, dim=-1) - temperatures, top_ks, top_ps = [], [], [] + temperatures, top_ks, top_ps, request_counts = [], [], [], [] for active_id, count in active_id_and_counts: sp = request_data[active_id]["sampling"] - temperatures += [sp["temperature"]] * count - top_ks += [sp["top_k"]] * count - top_ps += [sp["top_p"]] * count - device = logits_2d.device + temperatures.append(sp["temperature"]) + top_ks.append(sp["top_k"]) + top_ps.append(sp["top_p"]) + request_counts.append(count) + expected_context = SimpleNamespace( + total_request_count=len(active_id_and_counts), + paused_request_count=0, + active_request_metadata={ + "temperature": torch.tensor(temperatures, dtype=torch.float32), + "top_k": torch.tensor(top_ks, dtype=torch.long), + "top_p": torch.tensor(top_ps, dtype=torch.float32), + }, + ) + row_to_request = torch.arange(len(request_counts)).repeat_interleave( + torch.tensor(request_counts) + ) return sampling.log_probs_kernel( - logits_2d, - torch.tensor(temperatures, device=device, dtype=torch.float32), - torch.tensor(top_ks, device=device, dtype=torch.long), - torch.tensor(top_ps, device=device, dtype=torch.float32), + logits_2d, expected_context, token_to_request_index=row_to_request ) # Populate gpu_view for calculate_log_probs (which reads from gpu_view). @@ -1409,9 +1801,9 @@ def expected_log_probs(logits, active_id_and_counts): dynamic_context.initialize_attention_state() dynamic_context.transfer_bookkeeping_to_gpu() - # Generate new logits for the decode step. Now each request contributes 1 token. + # Generate a padded decode buffer where each active request contributes 1 token. decode_logits = torch.randn( - 1, num_active_requests, vocab_size, device='cuda', dtype=torch.float32 + 1, num_active_requests + 3, vocab_size, device='cuda', dtype=torch.bfloat16 ) decode_new_tokens = torch.randint(0, 100, (num_active_requests,), device='cuda').long() decode_log_probs, decode_log_probs_full = dynamic_context.calculate_log_probs( @@ -1420,7 +1812,9 @@ def expected_log_probs(logits, active_id_and_counts): # Verify the stored decode log probabilities decode_active = [(req_id, 1) for req_id in request_data] - expected_decode_full = expected_log_probs(decode_logits, decode_active) + expected_decode_full = expected_log_probs( + decode_logits[:, :num_active_requests], decode_active + ) assert torch.allclose(decode_log_probs_full, expected_decode_full, atol=1e-6) expected_decode_log_probs = expected_decode_full.to(torch.float32) diff --git a/tests/unit_tests/inference/contexts/test_dynamic_prefix_caching.py b/tests/unit_tests/inference/contexts/test_dynamic_prefix_caching.py index 84898db60d8..fa166af9669 100644 --- a/tests/unit_tests/inference/contexts/test_dynamic_prefix_caching.py +++ b/tests/unit_tests/inference/contexts/test_dynamic_prefix_caching.py @@ -36,7 +36,7 @@ def teardown_class(cls): Utils.destroy_model_parallel() @staticmethod - def _mamba_config(): + def _mamba_config(mamba_chunk_size=128): from megatron.core.inference.config import MambaInferenceStateConfig return MambaInferenceStateConfig( @@ -45,6 +45,7 @@ def _mamba_config(): ssm_states_shape=(4, 16), conv_states_dtype=torch.float32, ssm_states_dtype=torch.float32, + mamba_chunk_size=mamba_chunk_size, ) def _ctx( @@ -56,6 +57,7 @@ def _ctx( rounder=64, enable_prefix_caching=True, max_tokens=None, + max_requests=None, prefix_caching_eviction_policy=PrefixCachingEvictionPolicy.LRU, mamba_config=None, prefix_caching_mamba_gb=None, @@ -80,6 +82,7 @@ def _ctx( paused_buffer_size_gb=0.2 * buffer_size_gb, block_size_tokens=block_size_tokens, max_tokens=max_tokens, + max_requests=max_requests, mamba_inference_state_config=mamba_config, use_flashinfer_fused_rope=None, unified_memory_level=0, @@ -126,6 +129,7 @@ class _StubEngine(DynamicInferenceEngine): def __init__(self, context: DynamicInferenceContext, *, enable_chunked_prefill=False): self.context = context self.enable_chunked_prefill = enable_chunked_prefill + self.cuda_graph_all_prefills = False self._prefix_coordination_waits = 0 self._loop = asyncio.new_event_loop() self.waiting_request_ids: deque = deque() @@ -428,6 +432,113 @@ def test_ref_count_lru(self): for bid in active_blocks: assert alloc3.block_ref_counts[bid.item()].item() == 1 + @pytest.mark.internal + def test_add_request_full_cache_partial_hit_pins_matched_blocks(self): + """On a partial prefix hit against a FULL cache, the matched + blocks must be pinned before allocation so LRU eviction cannot reclaim + one of them for the new (non-matched) block. + + Scenario (mirrors the descendant-first LRU edge case): a cached chain + H0/S0 -> H1/S1 (older) plus an unrelated cached root HX/SX (newer) fill + the pool. An incoming prompt H0 -> H1 -> H2 matches [S0, S1] and needs one + new block. If S0/S1 are not pinned first, descendant-first LRU evicts the + older leaf S1 and immediately reuses it, yielding block_table [S0, S1, S1] + and a dangling H2 -> missing H1 chain. Correct behavior evicts SX and + yields [S0, S1, SX] with a contiguous H0 -> H1 -> H2 chain. + """ + ctx = self._ctx() + bs = ctx.block_size_tokens + alloc = ctx.kv_block_allocator + + # Cached chain H0/S0 -> H1/S1, seeded with an OLD timestamp. + ctx.prefix_cache_lru_clock = 1 + req_chain = self._req(ctx, self._prompt(bs * 2)) + ctx.add_request(req_chain) + s0, s1 = self._block_ids(ctx, 0, 2) + h0, h1 = req_chain.precomputed_block_hashes[0], req_chain.precomputed_block_hashes[1] + + # Unrelated cached root HX/SX, seeded with a NEWER timestamp, so a naive + # oldest-first / descendant-first eviction would prefer the chain leaf. + ctx.prefix_cache_lru_clock = 10 + ctx.add_request(self._req(ctx, self._prompt(bs, offset=9000), request_id=2)) + (sx,) = self._block_ids(ctx, 1, 1) + + # All three slots are distinct and now cached (ref_count drops to 0). + assert len({s0, s1, sx}) == 3 + ctx.release_memory_blocks_from_request_indexes(torch.tensor([0, 1])) + ctx.total_request_count = 0 + assert alloc.block_ref_counts[s0].item() == 0 + assert alloc.block_ref_counts[s1].item() == 0 + assert alloc.block_ref_counts[sx].item() == 0 + + # Force a full pool: the new block for H2 can only come from eviction. + alloc.total_avail = 0 + + # Incoming prompt H0 -> H1 -> H2: first two blocks match the cached chain, + # the third (H2) is new and must trigger a single eviction. + ctx.prefix_cache_lru_clock = 20 + req_new = self._req(ctx, self._prompt(bs * 3), request_id=3) + ctx.add_request(req_new) + h2 = req_new.precomputed_block_hashes[2] + + block_table = self._block_ids(ctx, 0, 3) + + # Matched blocks are preserved and SX (the unrelated root) is evicted/reused. + assert block_table == [s0, s1, sx] + # All three block IDs are distinct — no duplicate from a reclaimed match. + assert len(set(block_table)) == 3 + # Matched blocks stay pinned for the new request. + assert alloc.block_ref_counts[s0].item() == 1 + assert alloc.block_ref_counts[s1].item() == 1 + assert alloc.block_ref_counts[sx].item() == 1 + # Contiguous H0 -> H1 -> H2 hash chain over [S0, S1, SX]. + assert alloc.block_hashes[s0].item() == h0 + assert alloc.block_hashes[s1].item() == h1 + assert alloc.block_hashes[sx].item() == h2 + # Parent bookkeeping is stored as resolved block ids: S1's parent is S0 + # and SX's parent is S1 along the H0 -> H1 -> H2 chain. + assert alloc.block_parent_id[s1].item() == s0 + assert alloc.block_parent_id[sx].item() == s1 + assert alloc.kv_hash_to_block_id[h1] == s1 + + @pytest.mark.internal + def test_check_availability_excludes_already_pinned_matches(self): + """check_availability reserves only matched blocks that are currently + evictable (ref_count == 0). A matched prefix already pinned by an + in-flight request frees no capacity when re-pinned, so reserving it would + under-report availability and needlessly defer shared-prefix requests.""" + ctx = self._ctx() + bs = ctx.block_size_tokens + alloc = ctx.kv_block_allocator + + # Request A stays active, pinning the shared prefix H0/S0 -> H1/S1. + ctx.add_request(self._req(ctx, self._prompt(bs * 2))) + s0, s1 = self._block_ids(ctx, 0, 2) + assert alloc.block_ref_counts[s0].item() == 1 + assert alloc.block_ref_counts[s1].item() == 1 + + # One unrelated block is cached and evictable (ref_count == 0). + ctx.add_request(self._req(ctx, self._prompt(bs, offset=9000), request_id=2)) + (sx,) = self._block_ids(ctx, 1, 1) + ctx.release_memory_blocks_from_request_indexes(torch.tensor([1])) + assert alloc.block_ref_counts[sx].item() == 0 + assert int(alloc.get_evictable_block_count()) == 1 + + # Free pool exhausted: the one new block B needs (H2) can only come from + # evicting SX. The already-pinned matches S0/S1 must not be reserved. + alloc.total_avail = 0 + + # Request B shares H0/H1 with A and needs one new block for H2. + req_b = self._req(ctx, self._prompt(bs * 3), request_id=3) + matched, num_from_pool, *_ = ctx._compute_prefix_match(req_b, req_b.remaining_prompt_length) + assert matched == [s0, s1] + assert num_from_pool == 1 + + _, _, kv_cache_available = ctx.check_availability(req_b) + # SX (the sole evictable block) can satisfy H2; reserving the pinned + # matches would wrongly report the request as un-addable. + assert kv_cache_available is True + @pytest.mark.internal def test_ref_count_refzero(self): bs = 32 @@ -688,17 +799,13 @@ def test_mamba_cache_lifecycle(self): ctx6.add_request(self._req(ctx6, p6.clone())) msa6 = ctx6.mamba_slot_allocator self._mamba_allocate_and_register(ctx6, self._block_ids(ctx6, 0, 4)[:2]) - engine6 = _StubEngine(ctx6) - assert engine6._find_mamba_match_count(self._req(ctx6, p6.clone(), request_id=2)) == 2 + req6 = self._req(ctx6, p6.clone(), request_id=2) + assert ctx6._find_mamba_match_count(req6, 0, len(req6.precomputed_block_hashes)) == 2 # no match when no mamba hashes registered ctx7 = self._mctx() ctx7.add_request(self._req(ctx7, self._prompt(bs * 3))) - assert ( - _StubEngine(ctx7)._find_mamba_match_count( - self._req(ctx7, self._prompt(bs * 3), request_id=2) - ) - == 0 - ) + req7 = self._req(ctx7, self._prompt(bs * 3), request_id=2) + assert ctx7._find_mamba_match_count(req7, 0, len(req7.precomputed_block_hashes)) == 0 # allocate, free, re-allocate ctx8 = self._mctx() @@ -718,6 +825,35 @@ def test_mamba_cache_lifecycle(self): and ctx8.mamba_slot_allocator.has_state(bids8[0]) ) + @pytest.mark.internal + def test_hybrid_prefix_caching_without_mamba_budget_warns(self, caplog): + # Memory-only mode: prefix caching on a hybrid model without a Mamba cache + # budget is allowed (KV prefixes deduplicated for memory savings) but must + # warn that Mamba state caching and prefill skipping are disabled, and must + # not allocate a slot allocator. + import logging as _logging + + with caplog.at_level(_logging.WARNING): + ctx = self._ctx( + mamba_config=self._mamba_config(), + enable_prefix_caching=True, + prefix_caching_mamba_gb=None, + ) + assert ctx.is_hybrid_model + assert ctx.mamba_slot_allocator is None + assert "memory-only" in caplog.text + + @pytest.mark.internal + def test_mamba_cache_budget_too_small_raises(self): + # The CUDA-graph extraction scratch (sized to the per-step token-budget + # cap, max_mamba_intermediate_states_per_step) is reserved from + # prefix_caching_mamba_gb before the durable cache is sized. A budget too + # small to fit the scratch plus at least one durable slot is a hard + # configuration error, not a silent over-allocation (which previously + # could OOM at startup). + with pytest.raises(ValueError, match="prefix cache budget"): + self._mctx(prefix_caching_mamba_gb=1e-5) + @pytest.mark.internal def test_mamba_prefill_skip_and_zero_prefill(self): # mamba match limits prefill skip @@ -864,6 +1000,156 @@ def test_mamba_intermediate_offsets(self): torch.full_like(ctx5.mamba_slot_allocator.conv_states[layer, slot5], layer + 1.0), ) + @pytest.mark.internal + def test_max_intermediate_states_per_step_formula(self): + # The extraction buffers are sized by the tighter of two per-step bounds: + # token-based: ceil(max_tokens / block_size) + 1 + # request-based: MAX_INTERMEDIATE_OFFSETS_PER_REQUEST * max_requests + import math + + from megatron.core.inference.contexts.mamba_slot_allocator import ( + MAX_INTERMEDIATE_OFFSETS_PER_REQUEST, + ) + + def token_based(ctx): + return math.ceil(ctx.max_tokens / ctx.block_size_tokens) + + def request_based(ctx): + return MAX_INTERMEDIATE_OFFSETS_PER_REQUEST * ctx.max_requests + + # Token-limited regime: many requests, so the token budget is tighter. + ctx = self._mctx(block_size_tokens=256, max_tokens=2048) + assert ctx.max_requests >= 3 # ensure this regime is actually token-limited + expected = min(token_based(ctx), request_based(ctx)) + assert expected == token_based(ctx) # token bound wins here + assert ctx.max_mamba_intermediate_states_per_step == expected + # The single value is shared everywhere it's consumed. + assert ctx.mamba_slot_allocator.max_intermediate_count == expected + assert ctx.mamba_metadata.max_intermediate_count == expected + assert ctx.mamba_slot_allocator.intermediate_ssm_out.shape[1] == expected + + # Request-limited regime: few requests but a large token budget, so + # 3 * max_requests is the tighter bound. This is the case the token-only + # formula over-allocated for (e.g. 1 request + 16384 tokens once reserved + # 65 scratch slots a single request could never fill). + ctx2 = self._mctx(block_size_tokens=256, max_tokens=2048, max_requests=2) + expected2 = min(token_based(ctx2), request_based(ctx2)) + assert expected2 == request_based(ctx2) # request bound wins here + assert expected2 < token_based(ctx2) # ...and it is strictly tighter + assert ctx2.max_mamba_intermediate_states_per_step == expected2 + assert ctx2.mamba_slot_allocator.max_intermediate_count == expected2 + assert ctx2.mamba_metadata.max_intermediate_count == expected2 + assert ctx2.mamba_slot_allocator.intermediate_ssm_out.shape[1] == expected2 + + @pytest.mark.internal + def test_intermediate_count_bounded_by_token_budget(self): + # Claim: a single engine step emits at most max_tokens / block_size Mamba + # intermediate states, regardless of how many prefill requests it packs. + # Fill the token budget with fresh multi-block prefills and confirm the + # extracted count never exceeds the scratch buffer. + bs = 256 + ctx = self._mctx(block_size_tokens=bs, max_tokens=2048, max_sequence_length=4096) + budget = ctx.max_mamba_intermediate_states_per_step + + # Non-block-aligned 2.5-block prefills (each crosses a block boundary on a + # mamba-chunk multiple -> one intermediate offset). Distinct content so + # they never prefix-match one another. + per_req = bs * 2 + bs // 2 # 640 tokens + n = ctx.max_tokens // per_req + assert n >= 2 + for i in range(n): + ctx.add_request( + self._req(ctx, self._prompt(per_req, offset=i * 100000), request_id=i + 1) + ) + + # Drive the step's metadata computation (populates intermediate_count). + ctx.initialize_attention_state() + ctx.transfer_bookkeeping_to_gpu() + + md = ctx.mamba_metadata + # Extraction actually fired (guards against a silent no-op test)... + assert md.intermediate_count > 0 + # ...and the packed step never exceeds the token-budget bound. + assert md.intermediate_count <= budget + assert md.intermediate_count == sum(md.per_request_intermediate_counts) + + @pytest.mark.internal + def test_intermediate_count_fills_scratch_buffer(self): + # Reviewer follow-up: drive a single step that consumes nearly the whole + # scratch buffer (not just a few slots) and confirm it is never overrun. + # + # Realistic config: attention block_size=256, mamba_chunk_size=128. Since + # 256 is a multiple of 128, a 256-aligned block boundary is also a mamba + # chunk boundary (extractable); the mamba boundary at 128 is NOT a block + # boundary, so it is never a candidate. + # + # Each request is a fresh 257-token prompt (one token past a block): it + # crosses exactly one block boundary at token 256 -> exactly one + # intermediate offset, while consuming ~one block of tokens. Packing the + # token budget with these drives intermediate_count to max_tokens // 257, + # close to the token-budget bound -- filling nearly every scratch slot, + # unlike test_intermediate_count_bounded_by_token_budget (~1/3 of them). + import math + + bs = 256 + ctx = self._mctx(block_size_tokens=bs, max_tokens=2048, max_sequence_length=4096) + assert ctx.mamba_chunk_size == 128 # bs=256 is a multiple of mamba chunk 128 + budget = ctx.max_mamba_intermediate_states_per_step + + per_req = bs + 1 # 257: crosses the block boundary at bs=256 (a mamba-chunk multiple) + n = ctx.max_tokens // per_req + assert n >= 2 + assert n <= ctx.max_requests # all packed into a single step + for i in range(n): + ctx.add_request( + self._req(ctx, self._prompt(per_req, offset=i * 100000), request_id=i + 1) + ) + + ctx.initialize_attention_state() + ctx.transfer_bookkeeping_to_gpu() + + md = ctx.mamba_metadata + # Every request contributed exactly one intermediate offset... + assert md.intermediate_count == n + assert md.intermediate_count == sum(md.per_request_intermediate_counts) + # ...the step never overruns the scratch buffer... + assert md.intermediate_count <= budget + # ...and it fills all but a small, *derived* deficit: the block-spillover + # (per_req > bs, so fewer requests fit than there are blocks). n <= + # max_requests forces the token-based bound, so budget == ceil(max_tokens / bs). + assert budget == math.ceil(ctx.max_tokens / bs) + expected_unfilled = math.ceil(ctx.max_tokens / bs) - n + assert budget - md.intermediate_count == expected_unfilled + + @pytest.mark.internal + def test_intermediate_offsets_use_configured_mamba_chunk_size(self): + # Regression guard for the fix that reads mamba_chunk_size from the model + # config instead of hardcoding 128 in compute_and_store_offsets. + # + # With a 64-token mamba chunk and 64-token blocks, a 65-token prompt's + # block boundary at token 64 is a valid mamba-chunk multiple + # (64 % 64 == 0), so its state must be extracted and cached. The old + # hardcoded filter (64 % 128 != 0) would have wrongly skipped it, caching + # nothing and leaving no resume point for a later turn. + bs = 64 + ctx = self._mctx( + mamba_config=self._mamba_config(mamba_chunk_size=64), + block_size_tokens=bs, + max_sequence_length=512, + ) + assert ctx.mamba_chunk_size == 64 + msa = ctx.mamba_slot_allocator + + # Fresh 65-token prompt: crosses exactly the block boundary at token 64. + ctx.add_request(self._req(ctx, self._prompt(bs + 1))) + + count = msa._intermediate_counts_cpu[0].item() + # The boundary at token 64 was recorded -- would be 0 under the old + # hardcoded-128 filter, since 64 % 128 != 0. + assert count == 1 + offsets = msa._intermediate_offsets_cpu[0, :count].tolist() + assert offsets == [bs] + class TestMixedCachedAndFreshPrefill(PrefixCachingTestBase): @@ -1021,8 +1307,11 @@ def test_allocate_slots_batch(self): assert mixed_slots[0] == pre_slot assert msa3.free_count == free_before3 - 2 # only 2 new - # Eviction: exhaust free pool, verify eviction fires and returns valid slots - ctx4 = self._mctx(prefix_caching_mamba_gb=0.001) + # Eviction: exhaust free pool, verify eviction fires and returns valid slots. + # Budget must cover the CUDA-graph extraction scratch (3 * max_requests + # slots) plus the durable cache; a budget too small to fit the scratch now + # raises (see test_mamba_cache_budget_too_small_raises). + ctx4 = self._mctx(prefix_caching_mamba_gb=0.01) msa4 = ctx4.mamba_slot_allocator total_slots = msa4.max_slots ctx4.add_request(self._req(ctx4, self._prompt(bs * 4))) @@ -1339,3 +1628,130 @@ def test_routing_survives_prefix_match_lru(self): assert np.allclose(alloc.get_block_routing(b0), routing_b0) assert alloc.get_block_routing(b1) is not None assert np.allclose(alloc.get_block_routing(b1), routing_b1) + + +class TestPrefixCacheReuse(PrefixCachingTestBase): + """Cross-request prefix reuse on hybrid (Mamba) models: + + - reset(preserve_prefix_cache=True) keeps the cache; a plain reset() clears it. + - Per-context prefill token accounting (computed vs skipped). + - Mamba state is extracted for the last complete block of a multi-chunk prompt. + """ + + @pytest.mark.internal + def test_reset_preserves_prefix_cache_when_requested(self): + # LRU + prefix caching enabled: reset(preserve_prefix_cache=True) keeps the + # KV hash index (so an idle dummy_forward does not wipe cross-request reuse), + # while a plain reset() clears it. + ctx = self._ctx(enable_prefix_caching=True) + bs = ctx.block_size_tokens + ctx.add_request(self._req(ctx, self._prompt(bs * 2))) + cached = dict(ctx.kv_block_allocator.kv_hash_to_block_id) + assert len(cached) == 2 + + ctx.reset(preserve_prefix_cache=True) + assert ctx.kv_block_allocator.kv_hash_to_block_id == cached # preserved + + ctx.reset() # default: full reset + assert len(ctx.kv_block_allocator.kv_hash_to_block_id) == 0 # cleared + + @pytest.mark.internal + @pytest.mark.parametrize("enable_prefix_caching", [False, True]) + @pytest.mark.parametrize("preserve_prefix_cache", [False, True]) + @pytest.mark.parametrize("preserve_counters", [False, True]) + def test_reset_counter_preservation_is_explicit( + self, enable_prefix_caching, preserve_prefix_cache, preserve_counters + ): + """Counter preservation is independent of prefix-cache configuration.""" + ctx = self._ctx(buffer_size_gb=0.01, rounder=8, enable_prefix_caching=enable_prefix_caching) + counter_values = { + "step_count": 3, + "prefix_cache_lru_clock": 4, + "lifetime_prefill_token_count": 5, + "async_sched_step_count": 6, + "async_sched_compaction_step_count": 7, + } + for name, value in counter_values.items(): + setattr(ctx, name, value) + ctx.total_request_count = 1 + ctx.active_token_count = 1 + ctx.request_ids[0] = 10 + + ctx.reset(preserve_prefix_cache=preserve_prefix_cache, preserve_counters=preserve_counters) + + expected_counters = ( + counter_values if preserve_counters else dict.fromkeys(counter_values, 0) + ) + assert {name: getattr(ctx, name) for name in counter_values} == expected_counters + assert ctx.total_request_count == 0 + assert ctx.active_token_count == 0 + assert ctx.request_ids[0] == -1 + + @pytest.mark.internal + def test_prefill_computed_and_skipped_counters(self): + # A second request that shares a cached prefix should skip that prefix's + # prefill; the per-context counters must reflect computed vs skipped tokens. + ctx = self._ctx(enable_prefix_caching=True) + bs = ctx.block_size_tokens + + ctx.add_request(self._req(ctx, self._prompt(bs * 4), request_id=1)) + assert ctx.prefix_cache_prefill_skipped_tokens == 0 + assert ctx.prefix_cache_prefill_computed_tokens == bs * 4 + + # request 2 shares the first 4 blocks, adds 2 new blocks + req2 = self._req(ctx, self._prompt(bs * 6), request_id=2) + (matched, _, _, _, prefix_skip, _) = ctx._compute_prefix_match(req2, bs * 6) + assert len(matched) == 4 and prefix_skip == bs * 4 + ctx.add_request(req2) + + assert ctx.prefix_cache_prefill_skipped_tokens == bs * 4 + assert ctx.prefix_cache_prefill_computed_tokens == bs * 6 # 4bs + 2bs + + @pytest.mark.internal + def test_mamba_extraction_covers_last_block_of_continuation_chunk(self): + # For a non-block-aligned, multi-chunk prompt, the last complete block lies + # in a continuation chunk. Extraction offsets are chunk-relative, so that + # boundary's Mamba state is recorded when its chunk is scheduled. + ctx = self._ctx( + mamba_config=self._mamba_config(), + prefix_caching_mamba_gb=0.01, + block_size_tokens=256, + max_sequence_length=4096, + ) # mamba prefix caching enabled + bs = ctx.block_size_tokens + assert bs == 256 + msa = ctx.mamba_slot_allocator + + prompt_len = bs * 3 + 64 # 3 complete blocks + a 64-token remainder + req = self._req(ctx, self._prompt(prompt_len)) + ctx.add_request(req) # populates request_to_kv_block_ids[0] + overall_blocks = ctx.request_kv_block_counts[0].item() + assert overall_blocks == 4 # ceil(832 / 256) + + # Simulate the continuation chunk that covers tokens [2*bs, prompt_len): + # finished=2*bs, no prefix skip, the rest of the prompt as the chunk. + req.finished_chunk_token_count = 2 * bs + cont_chunk = prompt_len - 2 * bs + msa.compute_and_store_offsets( + req, + current_id=0, + skip_tokens=0, + prefill_chunk_length=cont_chunk, + num_matched_blocks=0, + matched_block_ids=[], + overall_required_blocks=overall_blocks, + ) + + # last complete block boundary = 3*bs (768); chunk-relative offset = 768-512=256. + last_aligned_abs = (prompt_len // bs) * bs + expected_offset = last_aligned_abs - 2 * bs + count = msa._intermediate_counts_cpu[0].item() + assert count >= 1 + recorded = msa._intermediate_offsets_cpu[0, :count].tolist() + assert expected_offset in recorded + # the recorded boundary maps to the last complete block (index 2) + idx = recorded.index(expected_offset) + assert ( + msa._intermediate_block_ids_cpu[0, idx].item() + == ctx.request_to_kv_block_ids[0][last_aligned_abs // bs - 1].item() + ) diff --git a/tests/unit_tests/inference/contexts/test_kv_block_allocator.py b/tests/unit_tests/inference/contexts/test_kv_block_allocator.py index 57b58230171..2cb552ee14f 100644 --- a/tests/unit_tests/inference/contexts/test_kv_block_allocator.py +++ b/tests/unit_tests/inference/contexts/test_kv_block_allocator.py @@ -19,6 +19,7 @@ def _make_context( total_request_count=0, request_kv_block_counts=None, request_to_kv_block_ids=None, + prefix_cache_lru_clock=0, ): """Build a minimal DynamicInferenceContext-like fake for the allocator.""" if request_kv_block_counts is None: @@ -30,6 +31,7 @@ def _make_context( total_request_count=total_request_count, request_kv_block_counts=request_kv_block_counts, request_to_kv_block_ids=request_to_kv_block_ids, + prefix_cache_lru_clock=prefix_cache_lru_clock, ) @@ -112,7 +114,9 @@ def test_block_usage_counts_no_prefix_caching( ) def test_prefix_caching_state_layout(policy, expect_timestamps): """Prefix-caching mode allocates block_hashes (initially -1) and ref_counts - (initially 0). LRU policy also allocates timestamps; REF_ZERO does not.""" + (initially 0). LRU policy also allocates timestamps and the persisted + prefix-forest bookkeeping (block_parent_id / block_child_count); REF_ZERO + does not.""" a = KVBlockAllocator( _make_context(), total_count=8, @@ -124,6 +128,11 @@ def test_prefix_caching_state_layout(policy, expect_timestamps): assert (a.block_ref_counts == 0).all().item() assert a.kv_hash_to_block_id == {} assert hasattr(a, "block_timestamps") is expect_timestamps + assert hasattr(a, "block_parent_id") is expect_timestamps + assert hasattr(a, "block_child_count") is expect_timestamps + if expect_timestamps: + assert (a.block_parent_id == -1).all().item() + assert (a.block_child_count == 0).all().item() def test_prefix_caching_allocate_and_hash_registration(): @@ -143,15 +152,25 @@ def test_prefix_caching_allocate_and_hash_registration(): ids = a.allocate_memory_blocks(2) assert (a.block_ref_counts[ids] == 1).all().item() - # Hash registration populates both the tensor and the dict. + # Hash registration populates both the tensor and the dict. Parent hashes are + # ignored under REF_ZERO (they only drive LRU eviction ordering), so this mode + # keeps no per-block parent bookkeeping. a.register_kv_block_hashes(block_ids=[1, 3], block_hashes=[111, 333]) assert a.block_hashes[1].item() == 111 assert a.block_hashes[3].item() == 333 + assert not hasattr(a, "block_parent_id") assert a.kv_hash_to_block_id == {111: 1, 333: 3} + # Supplying parent hashes is accepted (and ignored) under REF_ZERO. + a.register_kv_block_hashes(block_ids=[2, 4], block_hashes=[222, 444], parent_hashes=[111, 222]) + + # Mismatched parent-hash length is rejected regardless of policy. + with pytest.raises(AssertionError): + a.register_kv_block_hashes(block_ids=[5], block_hashes=[555], parent_hashes=[1, 2]) + # Empty inputs are a no-op (avoids zero-element tensor construction). a.register_kv_block_hashes(block_ids=[], block_hashes=[]) - assert a.kv_hash_to_block_id == {111: 1, 333: 3} + assert a.kv_hash_to_block_id == {111: 1, 333: 3, 222: 2, 444: 4} # REF_ZERO has no eviction path when the free pool is short. small = KVBlockAllocator( @@ -189,3 +208,371 @@ def test_block_usage_counts_with_prefix_caching( a = KVBlockAllocator(ctx, total_count=TOTAL_COUNT, paused_count=3, enable_prefix_caching=True) assert a.get_active_used() == expected_active assert a.get_paused_used() == expected_paused + + +def test_release_shared_block_decrements_once_per_owner(): + """A shared prefix block appears once per finishing owner in a batched + release: each occurrence must decrement (scatter-accumulate), and a block + reaching ref 0 with a duplicated ID is freed/deregistered exactly once.""" + # REF_ZERO: three owners of a shared block finish in stages, with a private + # block mixed into the final batch. + a = KVBlockAllocator( + _make_context(), + total_count=8, + paused_count=2, + enable_prefix_caching=True, + prefix_caching_eviction_policy=PrefixCachingEvictionPolicy.REF_ZERO, + ) + ids = a.allocate_memory_blocks(2) # ref_count == 1 each + shared, private = int(ids[0]), int(ids[1]) + a.register_kv_block_hashes(block_ids=[shared], block_hashes=[111]) + a.block_ref_counts[shared] += 2 # two more owners pin the shared block -> ref 3 + avail0 = a.total_avail + + # One owner finishes alone: ref 3 -> 2, nothing freed yet. + a.release_memory_blocks(torch.tensor([shared], dtype=torch.int32)) + assert a.block_ref_counts[shared].item() == 2 + assert a.total_avail == avail0 + assert 111 in a.kv_hash_to_block_id + + # The final two owners and the private request finish in one batch: the + # shared block appears twice and both decrements must land (ref 2 -> 0). + a.release_memory_blocks(torch.tensor([shared, private, shared], dtype=torch.int32)) + assert a.block_ref_counts[shared].item() == 0 + assert a.block_ref_counts[private].item() == 0 + # Two distinct blocks return to the pool; the shared one only once (not twice). + assert a.total_avail == avail0 + 2 + assert 111 not in a.kv_hash_to_block_id # deregistered exactly once + free_region = a.block_bag[: a.total_avail].tolist() + assert len(set(free_region)) == len(free_region) # no double-returned id + + # LRU: a hashed shared block released by both owners in one batch must hit + # ref 0 (becoming evictable), not stall at 1 with a leaked reference. A hashed + # block stays cached for reuse rather than returning to the pool. + lru = KVBlockAllocator( + _make_context(), + total_count=8, + paused_count=2, + enable_prefix_caching=True, + prefix_caching_eviction_policy=PrefixCachingEvictionPolicy.LRU, + ) + lshared = int(lru.allocate_memory_blocks(1)[0]) + lru.register_kv_block_hashes(block_ids=[lshared], block_hashes=[333], parent_hashes=[0]) + lru.block_ref_counts[lshared] += 1 # second owner -> ref 2 + lru.release_memory_blocks(torch.tensor([lshared, lshared], dtype=torch.int32)) + assert lru.block_ref_counts[lshared].item() == 0 + assert int(lru.get_evictable_block_count()) == 1 + assert lru.block_hashes[lshared].item() == 333 # kept cached, not pool-returned + + +# --------------------------------------------------------------------------- +# LRU eviction: parent-chain safety +# --------------------------------------------------------------------------- + + +def _lru_allocator(total_count=16, paused_count=1): + """LRU-mode prefix-caching allocator over a fresh fake context.""" + return KVBlockAllocator( + _make_context(), + total_count=total_count, + paused_count=paused_count, + enable_prefix_caching=True, + prefix_caching_eviction_policy=PrefixCachingEvictionPolicy.LRU, + ) + + +def _seed_cached_chain(a, block_ids, hashes, parents, timestamps): + """Register a chain of cached (ref_count == 0) blocks with explicit LRU + timestamps, bypassing the allocation path to control the layout directly.""" + a.register_kv_block_hashes(block_ids=block_ids, block_hashes=hashes, parent_hashes=parents) + ids = torch.tensor(block_ids, dtype=torch.int64) + a.block_ref_counts[ids] = 0 # cached / evictable + a.block_timestamps[ids] = torch.tensor(timestamps, dtype=torch.int64) + # Mark the blocks as out of the free pool so _deregister_blocks (which pushes + # them back) keeps total_avail bookkeeping consistent. + a.total_avail -= len(block_ids) + + +def _assert_prefix_invariant(a): + """Every cached block must have its parent cached too (or be a root). This is + exactly the invariant _find_kv_match_count relies on.""" + cached_ids = set(a.kv_hash_to_block_id.values()) + for block_hash, block_id in a.kv_hash_to_block_id.items(): + parent_id = a.block_parent_id[block_id].item() + if parent_id >= 0: + assert parent_id in cached_ids, ( + f"dangling child: block {block_id} (hash {block_hash}) parent " + f"block {parent_id} not cached" + ) + + +def test_evict_lru_never_orphans_a_child(): + """Regression: with chunked prefill an ancestor block can end up OLDER than + its descendant. A naive oldest-first eviction would evict the parent and leave + a dangling child; leaf-only eviction must evict the child instead.""" + a = _lru_allocator() + # Chain b0 -> b1 -> b2. Parent b1 (ts=1) is older than child b2 (ts=5). + _seed_cached_chain( + a, block_ids=[0, 1, 2], hashes=[10, 20, 30], parents=[0, 10, 20], timestamps=[1, 1, 5] + ) + + assert a.evict_lru_blocks(1) is True + # The leaf (b2, hash 30) is evicted, not the older parent b1 (hash 20). + assert a.kv_hash_to_block_id == {10: 0, 20: 1} + assert a.block_hashes[2].item() == -1 + assert a.block_parent_id[2].item() == -1 + # Evicting the leaf drops it from its parent's child count. + assert a.block_child_count[1].item() == 0 + _assert_prefix_invariant(a) + + +def test_evict_lru_cascades_up_the_chain(): + """Evicting more blocks than there are leaves walks up the chain from the + deepest descendant, always keeping the retained set descendant-closed.""" + a = _lru_allocator() + _seed_cached_chain( + a, block_ids=[0, 1, 2], hashes=[10, 20, 30], parents=[0, 10, 20], timestamps=[1, 1, 5] + ) + + assert a.evict_lru_blocks(2) is True + # b2 then b1 evicted; only the root b0 remains. + assert a.kv_hash_to_block_id == {10: 0} + _assert_prefix_invariant(a) + + +def test_evict_lru_normal_lru_order_when_leaf_is_oldest(): + """When the oldest block is already a leaf (the common partial-match case, + where ancestors are refreshed and descendants are stale), plain LRU order + applies and the oldest leaf is evicted first.""" + a = _lru_allocator() + # Ancestors refreshed (ts=9); descendant stale (ts=3) and is the leaf. + _seed_cached_chain( + a, block_ids=[0, 1, 2], hashes=[10, 20, 30], parents=[0, 10, 20], timestamps=[9, 9, 3] + ) + + assert a.evict_lru_blocks(1) is True + assert a.kv_hash_to_block_id == {10: 0, 20: 1} + _assert_prefix_invariant(a) + + +def test_evict_lru_branching_prefix_tree(): + """A shared parent with two divergent children (branching prefixes) must keep + the parent cached until BOTH children are evicted.""" + a = _lru_allocator() + # b0 is the parent of both b1 and b2 (e.g. prompts "P+X" and "P+Y"). + _seed_cached_chain( + a, block_ids=[0, 1, 2], hashes=[10, 20, 30], parents=[0, 10, 10], timestamps=[1, 2, 8] + ) + + # Evicting one block takes a leaf (b1, the older child), never the parent. + assert a.evict_lru_blocks(1) is True + assert a.kv_hash_to_block_id == {10: 0, 30: 2} + _assert_prefix_invariant(a) + + # Evicting the second child leaves only the parent. + assert a.evict_lru_blocks(1) is True + assert a.kv_hash_to_block_id == {10: 0} + _assert_prefix_invariant(a) + + +def test_evict_lru_cached_child_with_pinned_parent_treated_as_root(): + """Multi-turn / agentic case: a shared prefix block S0 stays pinned by an + active request (ref_count > 0) while a descendant S1 from a finished turn is + cached (ref_count == 0). S0 is not in the candidate set, so S1's parent + resolves to -1 (S1 is treated as a forest root) and is safely evicted by + normal LRU. The pinned parent must never be touched, even when it is the + oldest block of all.""" + a = _lru_allocator() + # Chain S0 -> S1, plus an unrelated cached root SX. + a.register_kv_block_hashes( + block_ids=[0, 1, 2], block_hashes=[10, 20, 30], parent_hashes=[0, 10, 0] + ) + ids = torch.tensor([0, 1, 2], dtype=torch.int64) + # S0 pinned (active request), S1 and SX cached/evictable. S0 is the OLDEST + # (ts=0) — a pin-blind oldest-first eviction would wrongly take it and orphan + # nothing here, but in general orphan its children. + a.block_ref_counts[ids] = torch.tensor([1, 0, 0], dtype=torch.int32) + a.block_timestamps[ids] = torch.tensor([0, 1, 9], dtype=torch.int64) + a.total_avail -= 3 + + # Only S1 and SX are candidates; the pinned S0 is excluded. + assert int(a.get_evictable_block_count()) == 2 + + # Evict one: S1 (ts=1) is the oldest candidate and a leaf; evicted first. + assert a.evict_lru_blocks(1) is True + assert a.kv_hash_to_block_id == {10: 0, 30: 2} # S0 (pinned) + SX survive + assert a.block_ref_counts[0].item() == 1 # parent still pinned + assert a.block_hashes[0].item() == 10 # parent hash intact + assert a.block_hashes[1].item() == -1 # child deregistered + _assert_prefix_invariant(a) + + # Evict again: only SX remains as a candidate; S0 stays pinned throughout. + assert a.evict_lru_blocks(1) is True + assert a.kv_hash_to_block_id == {10: 0} + assert a.block_ref_counts[0].item() == 1 + # The pinned parent can never be evicted, so a third eviction fails. + assert a.evict_lru_blocks(1) is False + + +def test_evict_lru_partial_chain_eviction_peels_from_leaf_keeping_root(): + """Evicting fewer blocks than a chain's length peels from the leaf end, even + when the root is the least-recently-used block. + + Chain A -> B -> C with the root A oldest (ts 1 < 2 < 3); evict 2. Eviction + proceeds leaf-first, so C then B are removed and the root A is retained. The + retained cache stays descendant-closed (no cached block is left with an + evicted parent). + """ + a = _lru_allocator() + _seed_cached_chain( + a, block_ids=[0, 1, 2], hashes=[10, 20, 30], parents=[0, 10, 20], timestamps=[1, 2, 3] + ) + + assert a.evict_lru_blocks(2) is True + # Leaf C and its parent B are evicted; the root A survives despite being oldest. + assert a.kv_hash_to_block_id == {10: 0} + assert a.block_hashes[0].item() == 10 # root A retained + assert a.block_hashes[1].item() == -1 # B deregistered + assert a.block_hashes[2].item() == -1 # C deregistered + _assert_prefix_invariant(a) + + +def test_evict_lru_insufficient_cached_blocks_returns_false(): + """When fewer cached blocks exist than requested, eviction fails without + touching the cache.""" + a = _lru_allocator() + _seed_cached_chain(a, block_ids=[0, 1], hashes=[10, 20], parents=[0, 10], timestamps=[1, 2]) + assert a.evict_lru_blocks(3) is False + assert a.kv_hash_to_block_id == {10: 0, 20: 1} + + +def test_evict_lru_keeps_hottest_leaf_over_cold_interior_parent(): + """Optimality: leaf-peeling must retain the single most-recently-used block + even when reaching it means evicting a colder interior parent elsewhere. A + block is kept only for its own recency, never because a hot descendant props + it up, so the hot leaf E survives while the colder interior block B is evicted. + + A(ts 1) -> B(ts 2) -> C(ts 5) + \-> F(ts 3) + \-> D(ts 3) -> E(ts 5) + """ + a = _lru_allocator(total_count=8) + # hashes: A=10, B=20, C=30, F=40, D=50, E=60 + _seed_cached_chain( + a, + block_ids=[0, 1, 2, 3, 4, 5], + hashes=[10, 20, 30, 40, 50, 60], + parents=[0, 10, 20, 20, 10, 50], + timestamps=[1, 2, 5, 3, 3, 5], + ) + + assert a.evict_lru_blocks(3) is True + # Evicted F(3), C(5), then B(2) once childless. Retains A, D, and the hottest + # block E -- never evicting E in favor of the colder interior B. + assert a.kv_hash_to_block_id == {10: 0, 50: 4, 60: 5} + assert a.block_hashes[5].item() == 60 # hottest leaf E retained + assert a.block_hashes[1].item() == -1 # cold interior B evicted + _assert_prefix_invariant(a) + + +def test_evict_lru_asserts_on_cyclic_parent_graph(): + """The parent graph is assumed acyclic (a forest). A hash collision producing + a cycle exposes no leaf, so the peel cannot collect enough blocks; this is a + bug and must fail loudly rather than silently under-evict.""" + a = _lru_allocator() + # 2-cycle: block 0's parent hash is 20 (block 1) and block 1's parent hash is + # 10 (block 0). register_kv_block_hashes never produces this — we seed it + # directly to model the pathological collision case. + _seed_cached_chain(a, block_ids=[0, 1], hashes=[10, 20], parents=[20, 10], timestamps=[1, 2]) + assert int(a.get_evictable_block_count()) == 2 + + with pytest.raises(AssertionError): + a.evict_lru_blocks(1) + + +def test_is_memory_available_excludes_soon_to_be_pinned_blocks(): + """potential_matched_count removes soon-to-be-pinned cached blocks from the + evictable capacity, so availability matches what allocation can satisfy + once those blocks (e.g. prefix matches) are pinned.""" + a = _lru_allocator(total_count=6, paused_count=1) + # Drain the free pool: every block is allocated (ref_count == 1), none free. + a.allocate_memory_blocks(a.total_avail) + assert a.total_avail == 0 + # Mark two blocks as cached/evictable, mirroring an LRU release: ref_count + # drops to 0 and the hash is retained, but the block stays out of the free + # pool (total_avail unchanged). + a.register_kv_block_hashes(block_ids=[0, 1], block_hashes=[10, 20], parent_hashes=[0, 10]) + a.block_ref_counts[torch.tensor([0, 1])] = 0 + assert a.total_avail == 0 + assert int(a.get_evictable_block_count()) == 2 + + # Both evictable blocks count toward availability by default. + assert a.is_memory_available(2) is True + # Excluding one (it will be pinned) leaves only one usable for the request. + assert a.is_memory_available(2, potential_matched_count=1) is False + assert a.is_memory_available(1, potential_matched_count=1) is True + # Excluding all evictable blocks leaves nothing to satisfy a new block. + assert a.is_memory_available(1, potential_matched_count=2) is False + + +def _reference_leaf_peel(block_ids, hashes, parents, timestamps, k_evict): + """Independent, straightforward greedy reference: repeatedly evict the + currently-evictable leaf with the oldest (timestamp, block_id). Returns the + set of evicted block ids. Used to pin the optimal eviction choice.""" + hash_to_id = dict(zip(hashes, block_ids)) + ts = dict(zip(block_ids, timestamps)) + child_count = {b: 0 for b in block_ids} + parent_of = {} + for b, p in zip(block_ids, parents): + pid = hash_to_id.get(p) + parent_of[b] = pid + if pid is not None: + child_count[pid] += 1 + + import heapq as _heapq + + heap = [(ts[b], b) for b in block_ids if child_count[b] == 0] + _heapq.heapify(heap) + evicted = set() + while heap and len(evicted) < k_evict: + _, b = _heapq.heappop(heap) + evicted.add(b) + pid = parent_of[b] + if pid is not None: + child_count[pid] -= 1 + if child_count[pid] == 0: + _heapq.heappush(heap, (ts[pid], pid)) + return evicted + + +def test_evict_lru_preserves_invariant_under_random_chains(): + """Property test: across many randomized multi-chain layouts and eviction + counts, eviction (a) preserves the parent-chain invariant and (b) evicts + exactly the optimal leaf-peel set (matched against an independent + reference).""" + torch.manual_seed(0) + for _ in range(50): + n = int(torch.randint(2, 10, (1,)).item()) + a = _lru_allocator(total_count=n + 4) + block_ids = list(range(n)) + # Build a forest: block k's parent is a random earlier block or a root. + hashes = [100 + k for k in range(n)] + parents = [] + for k in range(n): + if k == 0 or int(torch.randint(0, 2, (1,)).item()) == 0: + parents.append(0) # root + else: + parents.append(hashes[int(torch.randint(0, k, (1,)).item())]) + # Distinct timestamps so the optimal evicted set is unique and the + # reference comparison is exact (no tie-break ambiguity). + timestamps = torch.randperm(50)[:n].add(1).tolist() + _seed_cached_chain(a, block_ids, hashes, parents, timestamps) + + k_evict = int(torch.randint(1, n + 1, (1,)).item()) + expected_evicted = _reference_leaf_peel(block_ids, hashes, parents, timestamps, k_evict) + + assert a.evict_lru_blocks(k_evict) is True + retained = set(a.kv_hash_to_block_id.values()) + assert retained == set(block_ids) - expected_evicted + assert len(retained) == n - k_evict + _assert_prefix_invariant(a) diff --git a/tests/unit_tests/inference/coordinator_test_utils.py b/tests/unit_tests/inference/coordinator_test_utils.py index d33d8790ef9..0586231abc7 100644 --- a/tests/unit_tests/inference/coordinator_test_utils.py +++ b/tests/unit_tests/inference/coordinator_test_utils.py @@ -53,6 +53,7 @@ def make_coordinator_direct( coordinator.identities_of_data_parallel_ranks = deque( [rank_name_template.format(i).encode() for i in range(data_parallel_size)] ) + coordinator.removed_engine_identities = set() if deterministic_mode: coordinator.identities_of_data_parallel_ranks = deque( sorted(coordinator.identities_of_data_parallel_ranks) @@ -64,7 +65,6 @@ def make_coordinator_direct( n_ranks = data_parallel_size coordinator._hash_table = {} coordinator._hash_assignment_counter = 0 - coordinator._round_robin_idx = 0 sorted_identities = sorted(coordinator.identities_of_data_parallel_ranks) coordinator.identity_to_rank_index = { diff --git a/tests/unit_tests/inference/engines/test_cg_admission_gating.py b/tests/unit_tests/inference/engines/test_cg_admission_gating.py new file mode 100644 index 00000000000..74d061fd640 --- /dev/null +++ b/tests/unit_tests/inference/engines/test_cg_admission_gating.py @@ -0,0 +1,550 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +""" Unit tests for CUDA-graph-aware admission gating. """ + +import logging +import types + +import pytest + +from megatron.core.inference.batch_dimensions_utils import InferenceBatchDimensions +from megatron.core.inference.engines.dynamic_engine import DynamicInferenceEngine + + +def _create_engine( + cg_list, active_tok=0, num_prefill=0, num_decode=0, is_hybrid=False, warn_after=100 +): + """Mock engine instance.""" + engine = types.SimpleNamespace() + engine.context = types.SimpleNamespace( + cuda_graph_batch_dimensions_list=cg_list, + active_token_count=active_tok, + num_prefill_requests=num_prefill, + num_decode_requests=num_decode, + is_hybrid_model=is_hybrid, + use_cuda_graphs_for_non_decode_steps=True, + ) + engine.cuda_graph_all_prefills = True + engine._cg_admission_warn_after = warn_after + engine._cg_admission_gating_active = DynamicInferenceEngine._cg_admission_gating_active.__get__( + engine + ) + engine._find_cg_chunk_size = DynamicInferenceEngine._find_cg_chunk_size.__get__(engine) + engine._matches_cg_admission = DynamicInferenceEngine._matches_cg_admission.__get__(engine) + engine._cg_admission_check = DynamicInferenceEngine._cg_admission_check.__get__(engine) + engine._register_cg_wait = DynamicInferenceEngine._register_cg_wait.__get__(engine) + return engine + + +def _make_request(request_id=1, cg_wait_iters=0): + """Tiny stand-in for DynamicInferenceRequest; gating only reads/writes these fields.""" + return types.SimpleNamespace(request_id=request_id, cg_wait_iters=cg_wait_iters) + + +def _get_cudagraph(token_count, p, d): + return InferenceBatchDimensions( + token_count=token_count, prefill_req_count=p, decode_req_count=d + ) + + +# CG list sorted descending by token_count, matching the production list ordering. +SAMPLE_CG_LIST = [ + _get_cudagraph(256, 1, 255), + _get_cudagraph(256, 4, 252), + _get_cudagraph(256, 256, 0), + _get_cudagraph(128, 1, 127), + _get_cudagraph(128, 4, 124), + _get_cudagraph(128, 128, 0), + _get_cudagraph(64, 1, 63), + _get_cudagraph(64, 4, 60), + _get_cudagraph(64, 64, 0), + _get_cudagraph(16, 1, 15), + _get_cudagraph(16, 4, 12), + _get_cudagraph(16, 16, 0), + _get_cudagraph(4, 1, 3), + _get_cudagraph(4, 4, 0), + _get_cudagraph(2, 1, 1), + _get_cudagraph(2, 2, 0), +] + + +class TestGatingActivation: + """Gating must be strictly opt-in via cuda_graph_all_prefills. + + Configs and tests that exercise the scheduler with use_cuda_graphs_for_non_decode_steps=True + but cuda_graph_all_prefills=False will see the original scheduler behavior with no admission + gating. + """ + + def test_inactive_when_all_prefills_off(self): + engine = _create_engine(SAMPLE_CG_LIST) + engine.cuda_graph_all_prefills = False + assert engine._cg_admission_gating_active() is False + + def test_inactive_when_no_non_decode_graphs(self): + engine = _create_engine(SAMPLE_CG_LIST) + engine.context.use_cuda_graphs_for_non_decode_steps = False + assert engine._cg_admission_gating_active() is False + + def test_inactive_when_cg_list_empty(self): + engine = _create_engine([]) + assert engine._cg_admission_gating_active() is False + + def test_active_when_all_three_conditions_hold(self): + engine = _create_engine(SAMPLE_CG_LIST) + assert engine._cg_admission_gating_active() is True + + +class TestFindCgChunkSize: + """_find_cg_chunk_size should snap to the largest CG-aligned chunk within budget.""" + + def test_picks_largest_chunk_in_budget(self): + # Empty active state, large budget — should pick the largest captured token_count. + engine = _create_engine(SAMPLE_CG_LIST, active_tok=0, num_prefill=0, num_decode=0) + assert engine._find_cg_chunk_size(max_chunk_tokens=500) == 256 + + def test_respects_budget_ceiling(self): + # Budget below largest CG — should pick the largest CG that still fits. + engine = _create_engine(SAMPLE_CG_LIST) + assert engine._find_cg_chunk_size(max_chunk_tokens=100) == 64 + assert engine._find_cg_chunk_size(max_chunk_tokens=20) == 16 + assert engine._find_cg_chunk_size(max_chunk_tokens=5) == 4 + + def test_accounts_for_active_tokens(self): + # Already 50 tokens in flight; chunk + active must land on a CG boundary. + engine = _create_engine(SAMPLE_CG_LIST, active_tok=50) + # Need cg.token_count - 50 in [1, max_chunk]. With max=300: + # 256 - 50 = 206 (fits, valid). 128 - 50 = 78. 64 - 50 = 14. ... + # Largest fitting: 256 → chunk = 206. + assert engine._find_cg_chunk_size(max_chunk_tokens=300) == 206 + + def test_returns_none_when_no_cg_fits(self): + # No captured CG has token_count in (active, active+max_chunk]; helper returns + # None so the caller can explicitly defer. Active=300, max_chunk=10 -> need + # cg.token_count in (300, 310], none exists. + engine = _create_engine(SAMPLE_CG_LIST, active_tok=300) + assert engine._find_cg_chunk_size(max_chunk_tokens=10) is None + + def test_strict_mode_filters_insufficient_decode(self): + # Hybrid model: matcher requires captured_D >= real_D. At active D=125 and + # adding 1 new prefill, candidate (X, 1, 125) needs captured D >= 125. + # Only (256, 1, 255) and (256, 4, 252) qualify on D; pick smallest token_count + # that fits in budget — both have token=256, so chunk=256 is returned. + engine = _create_engine( + SAMPLE_CG_LIST, active_tok=0, num_prefill=0, num_decode=125, is_hybrid=True + ) + assert engine._find_cg_chunk_size(max_chunk_tokens=300) == 256 + + def test_strict_mode_no_match_returns_none(self): + # Active D=200, only (256, *, 252) and (256, *, 255) have D >= 200, requiring + # chunk=256. With smaller budget no CG matches in strict mode. + engine = _create_engine( + SAMPLE_CG_LIST, active_tok=0, num_prefill=0, num_decode=200, is_hybrid=True + ) + assert engine._find_cg_chunk_size(max_chunk_tokens=100) is None + + def test_empty_cg_list_returns_none(self): + engine = _create_engine([], active_tok=0) + assert engine._find_cg_chunk_size(max_chunk_tokens=50) is None + + +class TestCgAdmissionCheck: + """`_cg_admission_check` returns admission decision and updates request state.""" + + def test_match_returns_true_and_resets_counter(self): + engine = _create_engine(SAMPLE_CG_LIST) + req = _make_request(cg_wait_iters=5) # was previously deferred + candidate = _get_cudagraph(64, 1, 0) + assert engine._cg_admission_check(req, candidate) is True + assert req.cg_wait_iters == 0 + + def test_no_match_returns_false_and_increments_counter(self): + engine = _create_engine([]) # no captured graphs at all + req = _make_request() + candidate = _get_cudagraph(64, 1, 0) + assert engine._cg_admission_check(req, candidate) is False + assert req.cg_wait_iters == 1 + + def test_repeated_misses_accumulate(self): + engine = _create_engine([]) + req = _make_request() + for expected in range(1, 6): + engine._cg_admission_check(req, _get_cudagraph(64, 1, 0)) + assert req.cg_wait_iters == expected + + def test_warning_fires_at_threshold(self, caplog): + engine = _create_engine([], warn_after=3) + req = _make_request() + with caplog.at_level(logging.WARNING): + for _ in range(3): + engine._cg_admission_check(req, _get_cudagraph(64, 1, 0)) + starvation_warnings = [ + r for r in caplog.records if "deferred by CG-aware admission" in r.message + ] + assert len(starvation_warnings) == 1 + assert "3 steps" in starvation_warnings[0].message + + def test_warning_does_not_fire_below_threshold(self, caplog): + engine = _create_engine([], warn_after=100) + req = _make_request() + with caplog.at_level(logging.WARNING): + for _ in range(99): + engine._cg_admission_check(req, _get_cudagraph(64, 1, 0)) + assert not any("deferred by CG-aware admission" in r.message for r in caplog.records) + + def test_warning_repeats_at_each_multiple(self, caplog): + engine = _create_engine([], warn_after=2) + req = _make_request() + with caplog.at_level(logging.WARNING): + for _ in range(6): + engine._cg_admission_check(req, _get_cudagraph(64, 1, 0)) + starvation_warnings = [ + r for r in caplog.records if "deferred by CG-aware admission" in r.message + ] + # Fires at cg_wait_iters = 2, 4, 6. + assert len(starvation_warnings) == 3 + + def test_strict_vs_non_strict_decode_spillover(self): + # CGs with high total slots but limited per-type D; only non-strict can absorb + # the extra decodes by repurposing prefill slots. + cg_list = [_get_cudagraph(128, 128, 0), _get_cudagraph(128, 64, 64)] + candidate = InferenceBatchDimensions( + token_count=64, prefill_req_count=1, decode_req_count=70 + ) + + # Strict: needs captured_D >= 70. (128,128,0).D=0 ✗, (128,64,64).D=64 ✗. No match. + strict_engine = _create_engine(cg_list, is_hybrid=True) + assert strict_engine._cg_admission_check(_make_request(), candidate) is False + + # Non-strict: total=128 >= 71 ✓ on either CG; both match. Admit. + non_strict_engine = _create_engine(cg_list, is_hybrid=False) + assert non_strict_engine._cg_admission_check(_make_request(), candidate) is True + + +class TestFindChunkSizeStrictBoundary: + """Regression coverage for the Mamba-at-max_requests strict-matching scenario.""" + + def test_strict_at_max_requests_finds_p_grid_match(self): + # P-grid {1, 2, 4, 8} captured; real wants P+1=3 with D=508 at max_requests=512. + # Strict matching needs captured P>=3 AND captured D>=508 — (4, 508) satisfies. + # Shows that with adequate P-grid coverage, strict admission at max_requests + # is feasible (the next-larger P value absorbs the new prefill). + cg_list = [ + _get_cudagraph(512, 1, 511), + _get_cudagraph(512, 2, 510), + _get_cudagraph(512, 4, 508), + _get_cudagraph(512, 8, 504), + ] + engine = _create_engine( + cg_list, active_tok=0, num_prefill=2, num_decode=508, is_hybrid=True + ) + chunk = engine._find_cg_chunk_size(max_chunk_tokens=512) + assert chunk == 512 + # Confirm admission check also succeeds for this candidate. + candidate = InferenceBatchDimensions( + token_count=512, prefill_req_count=3, decode_req_count=508 + ) + assert engine._cg_admission_check(_make_request(), candidate) is True + + def test_strict_above_max_decode_returns_no_match(self): + # Real (P=2, D=510). Adding 1 prefill → (P=3, D=510) total=513 exceeds max. + # No captured CG has D >= 510 except (512, 1, 511) which has P=1 < 3. + cg_list = [ + _get_cudagraph(512, 1, 511), + _get_cudagraph(512, 2, 510), + _get_cudagraph(512, 4, 508), + ] + engine = _create_engine( + cg_list, active_tok=0, num_prefill=2, num_decode=510, is_hybrid=True + ) + # Helper returns None to signal "no CG match" to the caller — explicit so the + # caller can't accidentally schedule an un-graphed batch. + assert engine._find_cg_chunk_size(max_chunk_tokens=512) is None + # and a subsequent admission check on the same candidate also fails. + req = _make_request() + candidate = InferenceBatchDimensions( + token_count=512, prefill_req_count=3, decode_req_count=510 + ) + assert engine._cg_admission_check(req, candidate) is False + + +# Captured set for the deferral-flow tests: P-grid {1, 2, 4, 8, max=512} with +# decode-only counterparts. Designed so candidates with specific P/D combos can +# either match (admit) or miss (defer) depending on engine state. +DEFERRAL_CG_LIST = [ + _get_cudagraph(512, 1, 511), + _get_cudagraph(512, 2, 510), + _get_cudagraph(512, 4, 508), + _get_cudagraph(512, 8, 504), + _get_cudagraph(256, 1, 255), + _get_cudagraph(256, 2, 254), + _get_cudagraph(256, 4, 252), + _get_cudagraph(256, 8, 248), + _get_cudagraph(64, 1, 63), + _get_cudagraph(64, 2, 62), + _get_cudagraph(64, 4, 60), + _get_cudagraph(64, 8, 56), + _get_cudagraph(8, 0, 8), + _get_cudagraph(64, 0, 64), + _get_cudagraph(256, 0, 256), + _get_cudagraph(512, 0, 512), +] + + +class TestSchedulerDeferralInteraction: + """Multi-call scenarios that exercise the deferral / resume flow. + + Validates three properties of CG-aware admission gating: + - When one request defers, another admittable request can still proceed. + - A deferred request gets admitted once state changes (e.g., a decode completes and active-D + drops). + - The deferral path never silently falls back to eager — `_cg_admission_check` + strictly returns False on miss, and no internal flag is flipped to "schedule eagerly anyway" + """ + + def test_admittable_request_proceeds_when_other_is_deferred(self): + # Engine state: active (P=2, D=510) total=512. Mamba strict. + # Captured (4, 508) has D=508 < 510, so a P=3 candidate (would defer) cannot match. + # But a pure-decode candidate (token=8, P=0, D=1) matches (8, 0, 8) — admits. + engine = _create_engine( + DEFERRAL_CG_LIST, active_tok=512, num_prefill=2, num_decode=510, is_hybrid=True + ) + + # Request A: a new prefill that would push P to 3 — no captured shape covers it + # in strict mode (no captured P>=3 AND D>=510). + req_a = _make_request(request_id=1) + candidate_a = InferenceBatchDimensions( + token_count=512, prefill_req_count=3, decode_req_count=510 + ) + assert engine._cg_admission_check(req_a, candidate_a) is False + assert req_a.cg_wait_iters == 1 + + # Request B: a decode-only candidate that does match a captured graph. + # Reset active state to a low-load scenario (admittable). In a real scheduler + # these are sequential admissions against an evolving state. + admit_engine = _create_engine( + DEFERRAL_CG_LIST, active_tok=0, num_prefill=0, num_decode=0, is_hybrid=True + ) + req_b = _make_request(request_id=2) + candidate_b = InferenceBatchDimensions( + token_count=8, prefill_req_count=0, decode_req_count=1 + ) + # Note: prefill_req_count=0 takes the decode-only branch in is_applicable_for_batch_dim, + # which checks captured_decode_req_count >= real_decode_req_count and captured P==0. + assert engine._cg_admission_check(req_b, candidate_b) is True + assert req_b.cg_wait_iters == 0 + + def test_deferred_request_admits_once_state_changes(self): + # Initial state: active (P=2, D=510). Candidate (P=3, D=510) misses in strict mode. + engine = _create_engine( + DEFERRAL_CG_LIST, active_tok=512, num_prefill=2, num_decode=510, is_hybrid=True + ) + req = _make_request() + candidate_high_d = InferenceBatchDimensions( + token_count=512, prefill_req_count=3, decode_req_count=510 + ) + # First admission attempt: defers. + assert engine._cg_admission_check(req, candidate_high_d) is False + assert req.cg_wait_iters == 1 + + # Second attempt with active state still at D=510: still defers, wait counter + # increments since request hasn't been admitted. + assert engine._cg_admission_check(req, candidate_high_d) is False + assert req.cg_wait_iters == 2 + + # Decodes complete: active D drops to 508. Now the candidate (P=3, D=508) + # fits within captured (4, 508) strictly: P=4>=3, D=508>=508, total>=511. + engine.context.num_decode_requests = 508 + engine.context.active_token_count = 510 + candidate_lower_d = InferenceBatchDimensions( + token_count=512, prefill_req_count=3, decode_req_count=508 + ) + assert engine._cg_admission_check(req, candidate_lower_d) is True + # Wait counter resets on successful admission — the deferred request was + # finally admitted in a subsequent scheduler pass. + assert req.cg_wait_iters == 0 + + def test_admission_helpers_never_signal_eager_fallback(self): + # The design invariant: on miss, `_cg_admission_check` returns False and + # `_register_cg_wait` bumps the counter — that's it. No flag is set, no + # alternate "schedule eagerly" path is taken. The scheduler's break-on-False + # contract is what preserves the "no eager fallback under cuda_graph_all_prefills" + # property. + engine = _create_engine([], active_tok=0, num_prefill=0, num_decode=0) + req = _make_request() + candidate = _get_cudagraph(64, 1, 0) + + before_state = ( + engine.context.active_token_count, + engine.context.num_prefill_requests, + engine.context.num_decode_requests, + ) + # Fire 20 consecutive misses to increment the wait counter and verify nothing else changed. + for _ in range(20): + assert engine._cg_admission_check(req, candidate) is False + after_state = ( + engine.context.active_token_count, + engine.context.num_prefill_requests, + engine.context.num_decode_requests, + ) + + # Engine state is untouched by the gating helpers; only the request's + # wait counter advances. This proves the helpers never bypass the deferral + # via some "go eager" side channel. + assert before_state == after_state + assert req.cg_wait_iters == 20 + # The request object only has the fields we explicitly track — no surprise + # "eager_fallback_armed" flag or similar appeared. + assert set(vars(req).keys()) == {"request_id", "cg_wait_iters"} + + def test_two_requests_progress_independently_across_iterations(self): + # Two distinct requests in the waiting queue. Request 1 misses (high D), + # Request 2 hits (decode-only). Across multiple scheduler iterations, + # request 2 makes progress on each iteration while request 1's wait + # counter accumulates — until decodes drop and request 1 unblocks too. + cg_list = DEFERRAL_CG_LIST + engine = _create_engine( + cg_list, active_tok=512, num_prefill=2, num_decode=510, is_hybrid=True + ) + req_blocked = _make_request(request_id=1) + req_admittable = _make_request(request_id=2) + + # Mismatching candidate: needs strict D>=510 with P>=3 -> no captured graph. + blocked_candidate = InferenceBatchDimensions( + token_count=512, prefill_req_count=3, decode_req_count=510 + ) + # Matching candidate for the admittable one (decode-only with covered D). + admittable_candidate = InferenceBatchDimensions( + token_count=8, prefill_req_count=0, decode_req_count=1 + ) + + # Simulate 3 scheduler iterations, each incrementing req_blocked's wait counter + # while req_admittable is admitted every step. + results = [] + for step in range(3): + blocked_admit = engine._cg_admission_check(req_blocked, blocked_candidate) + admittable_admit = engine._cg_admission_check(req_admittable, admittable_candidate) + results.append((blocked_admit, admittable_admit)) + + # The blocked one defers on every step; the admittable one admits every step (counter stays + # at 0 — it never accumulates because each step succeeds). + for blocked_admit, admittable_admit in results: + assert blocked_admit is False + assert admittable_admit is True + assert req_blocked.cg_wait_iters == 3 + assert req_admittable.cg_wait_iters == 0 + + # Now active D drops since a decode completed. Check the previously blocked request is + # admitted. + engine.context.num_decode_requests = 508 + engine.context.active_token_count = 510 + unblocked_candidate = InferenceBatchDimensions( + token_count=512, prefill_req_count=3, decode_req_count=508 + ) + assert engine._cg_admission_check(req_blocked, unblocked_candidate) is True + assert req_blocked.cg_wait_iters == 0 + + +_CHUNKED_PREFILL_CG_CASES = [ + # parameters are label, active_tok, num_prefill, num_decode, max_chunk, is_hybrid, + # is_continuing, expected_chunk + # - is_continuing=True -> gating is skipped entirely; result is max_chunk + # - CG match found -> result is the snapped CG-aligned) chunk + # - No CG match -> eager fallback: result is max_chunk, not a deferral + pytest.param( + # Fresh batch, large budget — gating active, CG match at 256. + 0, + 0, + 0, + 300, + False, + False, + 256, + id="new_request_cg_match", + ), + pytest.param( + # Budget below the smallest CG (min token_count=2) — no match. + # Chunked prefill falls back to eager: uses max_chunk=1, not deferred. + 0, + 0, + 0, + 1, + False, + False, + 1, + id="new_request_no_cg_match_eager_fallback", + ), + pytest.param( + # Continuing chunked prefill: gating is bypassed regardless of CG coverage. + # Expected result equals max_chunk. + 50, + 1, + 0, + 100, + False, + True, + 100, + id="continuing_chunked_prefill_gating_skipped", + ), +] + + +class TestChunkedPrefillCgGating: + """Parametrized coverage for the CG-gating decision inside schedule_chunked_prefill. + + Exercises three distinct paths: + - is_continuing=True : gating is skipped; result = max_chunk + - CG hit : result = snapped (CG-aligned) chunk size + - CG miss : eager fallback; result = max_chunk (not a deferral) + """ + + @pytest.mark.parametrize( + "active_tok,num_prefill,num_decode,max_chunk,is_hybrid,is_continuing,expected_chunk", + _CHUNKED_PREFILL_CG_CASES, + ) + def test_chunk_size_decision( + self, + active_tok, + num_prefill, + num_decode, + max_chunk, + is_hybrid, + is_continuing, + expected_chunk, + ): + engine = _create_engine( + SAMPLE_CG_LIST, + active_tok=active_tok, + num_prefill=num_prefill, + num_decode=num_decode, + is_hybrid=is_hybrid, + ) + + if engine._cg_admission_gating_active() and not is_continuing: + snapped = engine._find_cg_chunk_size(max_chunk) + chunk = snapped if snapped is not None else max_chunk + else: + # Gating skipped (is_continuing) or gating inactive. + chunk = max_chunk + + assert chunk == expected_chunk + + def test_no_cg_match_does_not_defer(self): + # Core invariant: when CG gating is active and no graph matches the budget, + # the chunked-prefill path uses max_chunk (eager), never defers. + # SAMPLE_CG_LIST smallest token_count=2; budget=1 guarantees no match. + engine = _create_engine(SAMPLE_CG_LIST, active_tok=0, num_prefill=0, num_decode=0) + result = engine._find_cg_chunk_size(max_chunk_tokens=1) + assert result is None # confirms the miss path + # Caller's eager fallback: chunk = max_chunk, not deferred. + chunk = result if result is not None else 1 + assert chunk == 1 + + def test_cg_match_resets_wait_counter(self): + # On a CG hit the wait counter must be reset to 0 (matches non-chunked behaviour). + engine = _create_engine(SAMPLE_CG_LIST, active_tok=0, num_prefill=0, num_decode=0) + req = _make_request(cg_wait_iters=7) + snapped = engine._find_cg_chunk_size(max_chunk_tokens=300) + assert snapped is not None # hit + req.cg_wait_iters = 0 # as the engine does on a hit + assert req.cg_wait_iters == 0 diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine.py b/tests/unit_tests/inference/engines/test_dynamic_engine.py index a7317c82949..3286fcb0845 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine.py @@ -18,6 +18,7 @@ from megatron.core import parallel_state from megatron.core.inference.config import ( + AsyncScheduleMode, InferenceConfig, KVCacheManagementMode, MambaInferenceStateConfig, @@ -125,6 +126,7 @@ class DynamicEngineTestConfig: fp8: bool = False model_provider: str = "gpt" return_log_probs: bool = False + logprobs_mode: str = "raw_logprobs" materialize_only_last_token_logits: bool = True skip_prompt_log_probs: bool = False enable_chunked_prefill: bool = False @@ -148,6 +150,10 @@ class DynamicEngineTestConfig: num_speculative_tokens: int = 0 position_embedding_type: str = "learned_absolute" sampling_backend: str = 'torch' + temperature: float = 1.0 + top_k: int = 0 + top_p: float = 0.0 + async_sched_mode: AsyncScheduleMode = AsyncScheduleMode.LEGACY # Sliding-window attention config. When `window_size` is None, SWA is # disabled and all layers do full causal attention. When set to a # `(left, right)` tuple, layers selected by `window_attn_skip_freq` use a @@ -246,6 +252,9 @@ def _build_requests(cls, test_config: DynamicEngineTestConfig) -> List[DynamicIn ), return_log_probs=test_config.return_log_probs, skip_prompt_log_probs=test_config.skip_prompt_log_probs, + temperature=test_config.temperature, + top_k=test_config.top_k, + top_p=test_config.top_p, ) if not hasattr(sampling_params, "num_tokens_total"): # Remove this if statement branch in megatron-core 0.16 @@ -307,6 +316,8 @@ def _build_inference_context( track_generated_token_events=test_config.track_generated_token_events, num_speculative_tokens=test_config.num_speculative_tokens, sampling_backend=test_config.sampling_backend, + async_sched_mode=test_config.async_sched_mode, + logprobs_mode=test_config.logprobs_mode, ), ) @@ -678,48 +689,56 @@ def test_simple(self, model_provider, num_cuda_graphs, inference_cuda_graph_scop for layer in model.decoder.layers: assert layer.cudagraph_manager.cudagraph_runners - # Validate generated tokens. - gpt_expected_generated_tokens = [ - [69, 85, 55, 74, 56, 89, 64, 59, 55, 67, 15, 58, 6, 37, 54, 47], - [29, 54, 33, 72, 45, 76, 41, 56, 28, 25, 17, 2, 61, 6, 98, 76], - [35, 78, 54, 16, 79, 98, 22, 5, 60, 0, 1, 76, 77, 11, 25, 7], - [25, 75, 57, 85, 81, 37, 88, 17, 71, 15, 70, 64, 50, 0, 64, 45], - [32, 5, 85, 75, 30, 68, 23, 33, 20, 26, 89, 20, 49, 28, 38, 81], - [33, 69, 32, 49, 93, 24, 33, 6, 54, 89, 92, 97, 42, 80, 50, 53], - [82, 78, 78, 65, 26, 5, 69, 36, 37, 99], - [51, 70, 22, 1, 87, 42, 36, 26, 27, 56, 82, 32, 8, 80, 20, 43], - ] - - mamba_expected_generated_tokens = [ - [69, 85, 55, 74, 85, 89, 64, 59, 55, 67, 15, 58, 6, 37, 34, 47], - [29, 16, 33, 30, 45, 76, 41, 46, 82, 17, 17, 2, 61, 6, 98, 76], - [35, 78, 54, 16, 79, 98, 22, 5, 37, 30, 1, 76, 5, 11, 25, 86], - [25, 75, 57, 85, 81, 59, 88, 38, 71, 15, 70, 64, 50, 0, 64, 45], - [32, 5, 85, 75, 30, 68, 23, 33, 20, 26, 35, 20, 49, 28, 34, 81], - [87, 69, 32, 49, 93, 24, 33, 6, 54, 89, 92, 97, 42, 80, 50, 53], - [82, 78, 78, 19, 70, 5, 97, 36, 37, 99], - [51, 70, 22, 1, 87, 42, 36, 26, 27, 56, 82, 32, 8, 20, 20, 43], - ] - - if model_provider == "gpt": - expected_generated_tokens_list = gpt_expected_generated_tokens - elif model_provider == "hybrid": - expected_generated_tokens_list = mamba_expected_generated_tokens - else: - raise ValueError(f"Invalid model_provider {model_provider}") + # Because the TextGenerationController produces different outputs on different DP ranks, + # only verify the accuracy of the output on DP rank 0. + if parallel_state.get_data_parallel_rank() == 0: + + # Validate generated tokens. + gpt_expected_generated_tokens = [ + [69, 85, 55, 74, 56, 89, 64, 59, 55, 67, 15, 58, 6, 37, 54, 47], + [29, 54, 33, 72, 45, 76, 41, 56, 28, 25, 17, 2, 61, 6, 98, 76], + [35, 78, 54, 16, 79, 98, 22, 5, 60, 0, 1, 76, 77, 11, 25, 7], + [25, 75, 57, 85, 81, 37, 88, 17, 71, 15, 70, 64, 50, 0, 64, 45], + [32, 5, 85, 75, 30, 68, 23, 33, 20, 26, 89, 20, 49, 28, 38, 81], + [33, 69, 32, 49, 93, 24, 33, 6, 54, 89, 92, 97, 42, 80, 50, 53], + [82, 78, 78, 65, 26, 5, 69, 36, 37, 99], + [51, 70, 22, 1, 87, 42, 36, 26, 27, 56, 82, 32, 8, 80, 20, 43], + ] + + mamba_expected_generated_tokens = [ + [69, 85, 55, 74, 85, 89, 64, 59, 55, 67, 15, 58, 6, 37, 34, 47], + [29, 16, 33, 30, 45, 76, 41, 46, 82, 17, 17, 2, 61, 6, 98, 76], + [35, 78, 54, 16, 79, 98, 22, 5, 37, 30, 1, 76, 5, 11, 25, 86], + [25, 75, 57, 85, 81, 59, 88, 38, 71, 15, 70, 64, 50, 0, 64, 45], + [32, 5, 85, 75, 30, 68, 23, 33, 20, 26, 35, 20, 49, 28, 34, 81], + [87, 69, 32, 49, 93, 24, 33, 6, 54, 89, 92, 97, 42, 80, 50, 53], + [82, 78, 78, 19, 70, 5, 97, 36, 37, 99], + [51, 70, 22, 1, 87, 42, 36, 26, 27, 56, 82, 32, 8, 20, 20, 43], + ] + + if model_provider == "gpt": + expected_generated_tokens_list = gpt_expected_generated_tokens + elif model_provider == "hybrid": + expected_generated_tokens_list = mamba_expected_generated_tokens + else: + raise ValueError(f"Invalid model_provider {model_provider}") - print(f"Validating {len(env.requests)} requests.") - print(f"Expected generated tokens: {expected_generated_tokens_list}") - print(f"Actual generated tokens: {[request.generated_tokens for request in env.requests]}") + print(f"Validating {len(env.requests)} requests.") + print(f"Expected generated tokens: {expected_generated_tokens_list}") + print( + f"Actual generated tokens: {[request.generated_tokens for request in env.requests]}" + ) - assert len(env.requests) == len(expected_generated_tokens_list) + assert len(env.requests) == len(expected_generated_tokens_list) - for request, expected_generated_tokens in zip(env.requests, expected_generated_tokens_list): - assert request.generated_tokens == expected_generated_tokens, ( - f"request {request.request_id}, " - f"result ({request.generated_tokens}) != " - f"expected ({expected_generated_tokens})." - ) + for request, expected_generated_tokens in zip( + env.requests, expected_generated_tokens_list + ): + assert request.generated_tokens == expected_generated_tokens, ( + f"request {request.request_id}, " + f"result ({request.generated_tokens}) != " + f"expected ({expected_generated_tokens})." + ) @pytest.mark.internal @pytest.mark.skipif( @@ -1088,6 +1107,370 @@ async def test_run_engine(self): engine_task.cancel() + @pytest.mark.internal + @pytest.mark.asyncio + @pytest.mark.skipif( + not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + ) + async def test_async_sched_run_engine_accepts_request_during_overlap(self): + """Verify async overlap yields so a new request can enter a running engine.""" + with torch.inference_mode(): + test_config = DynamicEngineTestConfig( + num_requests=2, + min_prompt_length=4, + max_prompt_length=4, + num_tokens_to_generate=16, + async_sched_mode=AsyncScheduleMode.ASYNC, + ) + env = self._build_test_env(test_config) + long_request, short_request = env.requests + for request in env.requests: + request.sampling_params.top_k = 1 + request.sampling_params.top_p = 0.0 + request.sampling_params.termination_id = -1 + short_request.sampling_params.num_tokens_to_generate = 2 + + engine_task = asyncio.create_task(env.engine.run_engine()) + try: + long_request_future = env.engine._add_request(long_request) + + while len(long_request.generated_tokens) < 2: + await asyncio.sleep(0) + + generated_count_at_submission = len(long_request.generated_tokens) + short_request_future = env.engine._add_request(short_request) + await asyncio.gather(long_request_future, short_request_future) + + assert generated_count_at_submission < 16 + assert len(short_request.generated_tokens) == 2 + finally: + engine_task.cancel() + await engine_task + + @pytest.mark.internal + @pytest.mark.skipif( + not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + ) + @pytest.mark.parametrize( + ("enable_chunked_prefill", "enable_prefix_caching"), + [(True, False), (False, True), (True, True)], + ) + @torch.inference_mode() + def test_async_sched_prefix_caching_and_chunked_prefill_e2e( + self, enable_chunked_prefill, enable_prefix_caching + ): + """Async output matches legacy for chunking, KV caching, and their combination.""" + + def run(mode): + test_config = DynamicEngineTestConfig( + num_requests=0, + num_tokens_to_generate=4, + max_sequence_length=768, + context_block_size_tokens=256, + context_max_tokens=384 if enable_chunked_prefill else 1024, + context_max_requests=4, + enable_chunked_prefill=enable_chunked_prefill, + enable_prefix_caching=enable_prefix_caching, + async_sched_mode=mode, + ) + env = self._build_test_env(test_config) + prompt = torch.arange(512, dtype=torch.int64, device="cuda") % ( + test_config.vocab_size - 1 + ) + outputs = {} + + def add_request(request_id): + env.engine.add_request( + request_id=request_id, + prompt=prompt.clone(), + sampling_params=SamplingParams( + num_tokens_to_generate=4, termination_id=-1, top_k=1, top_p=0.0 + ), + ) + + add_request(0) + env.engine.step_modern() + add_request(1) + while env.engine.has_unfinished_requests(): + result = env.engine.step_modern() + for record in result["finished_request_records"]: + request = record.merge() + outputs[request.request_id] = list(request.generated_tokens) + return env.engine, outputs + + _, legacy_outputs = run(AsyncScheduleMode.LEGACY) + async_engine, async_outputs = run(AsyncScheduleMode.ASYNC) + + assert async_outputs == legacy_outputs + assert all(len(tokens) == 4 for tokens in async_outputs.values()) + assert async_engine.context.async_sched_step_count > 0 + if enable_prefix_caching: + assert async_engine._prefill_tokens_skipped > 0 + + @pytest.mark.internal + @pytest.mark.skipif( + not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + ) + @pytest.mark.parametrize( + "feature_config, sampling_backend, temperature, top_k, top_p, num_cuda_graphs", + [ + pytest.param({}, "torch", 0.8, 10, 0.0, None, id="torch-eager-top-k"), + pytest.param({}, "torch", 1.0, 0, 0.9, 2, id="torch-graphed-forward-top-p"), + pytest.param({}, "flashinfer", 1.2, 0, 0.0, None, id="flashinfer-eager-unfiltered"), + pytest.param({}, "flashinfer", 0.8, 10, 0.9, 2, id="flashinfer-graphed-forward"), + pytest.param( + { + "model_provider": "hybrid", + "num_speculative_tokens": 1, + "num_requests": 2, + "num_tokens_to_generate": 4, + }, + "torch", + 0.8, + 8, + 0.0, + None, + id="mamba-mtp-top-k", + ), + ], + ) + @torch.inference_mode() + def test_async_sched_sampling_matches_legacy( + self, feature_config, sampling_backend, temperature, top_k, top_p, num_cuda_graphs + ): + """Require seeded sampling parity across scheduling modes. + + Args: + feature_config (dict): Additional cumulative feature configuration. + sampling_backend (str): Sampling implementation under test. + temperature (float): Sampling temperature used by every request. + top_k (int): Top-k filter used by every request. + top_p (float): Top-p filter used by every request. + num_cuda_graphs (Optional[int]): Number of CUDA graph buckets, or + `None` for eager execution. + """ + if sampling_backend == "flashinfer": + pytest.importorskip("flashinfer") + if feature_config.get("model_provider") == "hybrid": + skip_if_mamba_sequence_packing_not_available("hybrid") + + common_config = dict( + num_requests=4, + min_prompt_length=4, + max_prompt_length=4, + num_tokens_to_generate=6, + num_gap_steps=0, + use_fixed_output_lengths=True, + sampling_backend=sampling_backend, + temperature=temperature, + top_k=top_k, + top_p=top_p, + context_max_requests=8, + num_cuda_graphs=num_cuda_graphs, + force_build_cuda_graphs=num_cuda_graphs is not None, + use_cuda_graphs_for_non_decode_steps=False, + ) + common_config.update(feature_config) + generated_tokens = {} + final_env = None + for mode in (AsyncScheduleMode.LEGACY, AsyncScheduleMode.ASYNC): + final_env = self._run_test(async_sched_mode=mode, **common_config) + assert all(request.status == Status.COMPLETED for request in final_env.requests) + generated_tokens[mode] = [request.generated_tokens for request in final_env.requests] + + assert ( + generated_tokens[AsyncScheduleMode.ASYNC] == generated_tokens[AsyncScheduleMode.LEGACY] + ) + + @pytest.mark.internal + @pytest.mark.skipif( + not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + ) + @pytest.mark.parametrize( + "feature_config, sampling_backend, logprobs_mode, skip_prompt_log_probs, num_cuda_graphs", + [ + pytest.param({}, "torch", "raw_logprobs", False, None, id="torch-raw-prompt"), + pytest.param({}, "torch", "processed_logprobs", True, 2, id="torch-processed-graph"), + pytest.param({}, "flashinfer", "raw_logprobs", False, None, id="flashinfer-raw-prompt"), + pytest.param( + {}, "flashinfer", "processed_logprobs", True, 2, id="flashinfer-processed-graph" + ), + pytest.param( + { + "model_provider": "hybrid", + "num_speculative_tokens": 1, + "num_requests": 2, + "num_tokens_to_generate": 4, + }, + "torch", + "raw_logprobs", + True, + None, + id="mamba-mtp-raw", + ), + ], + ) + @torch.inference_mode() + def test_async_sched_log_probs_match_legacy( + self, + feature_config, + sampling_backend, + logprobs_mode, + skip_prompt_log_probs, + num_cuda_graphs, + ): + """Require prompt and generated logprob parity across scheduling modes. + + Args: + feature_config (dict): Additional cumulative feature configuration. + sampling_backend (str): Sampling implementation under test. + logprobs_mode (str): Raw or sampling-processed logprob mode. + skip_prompt_log_probs (bool): Whether to omit prompt logprobs. + num_cuda_graphs (Optional[int]): Number of CUDA graph buckets. + """ + if sampling_backend == "flashinfer": + pytest.importorskip("flashinfer") + if feature_config.get("model_provider") == "hybrid": + skip_if_mamba_sequence_packing_not_available("hybrid") + + common_config = dict( + num_requests=4, + min_prompt_length=4, + max_prompt_length=4, + num_tokens_to_generate=6, + num_gap_steps=0, + use_fixed_output_lengths=True, + model_provider="gpt", + sampling_backend=sampling_backend, + temperature=0.8, + top_k=8, + return_log_probs=True, + logprobs_mode=logprobs_mode, + materialize_only_last_token_logits=skip_prompt_log_probs, + skip_prompt_log_probs=skip_prompt_log_probs, + context_max_requests=8, + num_cuda_graphs=num_cuda_graphs, + force_build_cuda_graphs=num_cuda_graphs is not None, + use_cuda_graphs_for_non_decode_steps=False, + ) + common_config.update(feature_config) + outputs = {} + for mode in AsyncScheduleMode: + env = self._run_test(async_sched_mode=mode, **common_config) + outputs[mode] = [ + (request.generated_tokens, request.prompt_log_probs, request.generated_log_probs) + for request in env.requests + ] + + legacy_outputs = outputs[AsyncScheduleMode.LEGACY] + for legacy, actual in zip(legacy_outputs, outputs[AsyncScheduleMode.ASYNC]): + assert actual[0] == legacy[0] + assert (actual[1] or []) == pytest.approx(legacy[1] or []) + assert actual[2] == pytest.approx(legacy[2]) + + def _run_stop_word_schedule( + self, + test_config: DynamicEngineTestConfig, + stop_word: Optional[str] = None, + detokenize_stop_sequence: bool = False, + ) -> DynamicEngineTestEnv: + """Run a schedule where only the first request has a string stop word. + + Args: + test_config (DynamicEngineTestConfig): Engine configuration for the run. + stop_word (Optional[str]): Whitespace-delimited token IDs used as the stop word. + detokenize_stop_sequence (bool): Whether the completed output retains the stop word. + + Returns: + DynamicEngineTestEnv: Completed test environment. + """ + env = self._build_test_env(test_config) + env.engine.controller.tokenizer.bos = None + env.engine.controller.tokenizer.tokenize = lambda text: [ + int(token_id) for token_id in text.split() + ] + + for request_idx, request in enumerate(env.requests): + request.sampling_params.termination_id = -1 + request.sampling_params.detokenize_stop_sequence = detokenize_stop_sequence + if request_idx == 0 and stop_word is not None: + request.sampling_params.stop_words = [stop_word] + env.engine._add_request(request) + + while env.engine.has_unfinished_requests(): + env.engine.step_modern() + + return env + + @pytest.mark.internal + @pytest.mark.skipif( + not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + ) + @pytest.mark.parametrize( + "sampling_backend,num_cuda_graphs,detokenize_stop_sequence", + [ + pytest.param("torch", None, True, id="torch-eager-keep"), + pytest.param("torch", 2, False, id="torch-graph-strip"), + pytest.param("flashinfer", None, False, id="flashinfer-eager-strip"), + pytest.param("flashinfer", 2, True, id="flashinfer-graph-keep"), + ], + ) + @torch.inference_mode() + def test_async_sched_stop_words_match_legacy( + self, sampling_backend, num_cuda_graphs, detokenize_stop_sequence + ): + """Require string stop-word parity while survivor requests keep decoding. + + Args: + sampling_backend (str): Sampling backend under test. + num_cuda_graphs (Optional[int]): CUDA graph bucket count, or ``None`` for eager mode. + detokenize_stop_sequence (bool): Whether completed output retains the stop word. + """ + if sampling_backend == "flashinfer": + pytest.importorskip("flashinfer") + + common_config = dict( + num_requests=4, + min_prompt_length=4, + max_prompt_length=4, + num_tokens_to_generate=8, + num_gap_steps=0, + model_provider="gpt", + sampling_backend=sampling_backend, + temperature=1.0, + top_k=1, + context_max_requests=8, + num_cuda_graphs=num_cuda_graphs, + force_build_cuda_graphs=num_cuda_graphs is not None, + use_cuda_graphs_for_non_decode_steps=False, + ) + + probe_env = self._run_stop_word_schedule( + DynamicEngineTestConfig(async_sched_mode=AsyncScheduleMode.LEGACY, **common_config) + ) + stop_word_ids = probe_env.requests[0].generated_tokens[2:4] + stop_word = " ".join(str(token_id) for token_id in stop_word_ids) + + legacy_env = self._run_stop_word_schedule( + DynamicEngineTestConfig(async_sched_mode=AsyncScheduleMode.LEGACY, **common_config), + stop_word, + detokenize_stop_sequence, + ) + async_env = self._run_stop_word_schedule( + DynamicEngineTestConfig(async_sched_mode=AsyncScheduleMode.ASYNC, **common_config), + stop_word, + detokenize_stop_sequence, + ) + + legacy_tokens = [request.generated_tokens for request in legacy_env.requests] + async_tokens = [request.generated_tokens for request in async_env.requests] + assert async_tokens == legacy_tokens + assert len(async_tokens[0]) < common_config["num_tokens_to_generate"] + if detokenize_stop_sequence: + assert async_tokens[0][-len(stop_word_ids) :] == stop_word_ids + assert async_env.engine.context.async_sched_step_count > 0 + assert async_env.engine.context.async_sched_compaction_step_count > 0 + @pytest.mark.internal @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" @@ -1223,8 +1606,11 @@ def test_return_log_probs(self): @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) + @pytest.mark.parametrize("async_sched_mode", list(AsyncScheduleMode)) @torch.inference_mode() - def test_return_prompt_log_probs_with_zero_tokens_to_generate(self): + def test_return_prompt_log_probs_with_zero_tokens_to_generate( + self, async_sched_mode: AsyncScheduleMode + ): """Prompt log probs must be returned when scoring only (num_tokens_to_generate=0). Regression test for a prefill-step trimming bug: when a request generates @@ -1234,12 +1620,16 @@ def test_return_prompt_log_probs_with_zero_tokens_to_generate(self): sampled-token log prob at the tail). The fix trims the excess *trailing* log probs instead. This is the path exercised by loglikelihood / echo evaluations (e.g. lm-eval-harness sends ``max_tokens=0``). + + Args: + async_sched_mode (AsyncScheduleMode): Scheduling mode under test. """ env = self._run_test( return_log_probs=True, materialize_only_last_token_logits=False, skip_prompt_log_probs=False, num_tokens_to_generate=0, + async_sched_mode=async_sched_mode, ) validated_any = False @@ -2110,15 +2500,22 @@ def get_log_probs(chunked: bool, max_tokens: int): @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) + @pytest.mark.parametrize("async_sched_mode", list(AsyncScheduleMode)) @pytest.mark.parametrize("skip_prompt_log_probs", [True, False]) @torch.inference_mode() - def test_top_n_logprobs_dynamic(self, skip_prompt_log_probs: bool): - """ - Test that top_n_logprobs are computed correctly in dynamic batching mode. + def test_top_n_logprobs_dynamic( + self, skip_prompt_log_probs: bool, async_sched_mode: AsyncScheduleMode + ): + """Test that top_n_logprobs are computed correctly in dynamic batching mode. + Verifies: 1. top_n_logprobs are returned for generated tokens 2. skip_prompt_log_probs controls whether prompt top-n logprobs are skipped 3. The top-n values are consistent with the selected token's log prob + + Args: + skip_prompt_log_probs (bool): Whether to omit prompt top-n logprobs. + async_sched_mode (AsyncScheduleMode): Scheduling mode under test. """ # Build test environment with multiple requests of varying lengths test_config = DynamicEngineTestConfig( @@ -2127,6 +2524,7 @@ def test_top_n_logprobs_dynamic(self, skip_prompt_log_probs: bool): max_prompt_length=12, num_tokens_to_generate=4, materialize_only_last_token_logits=False, + async_sched_mode=async_sched_mode, ) env = self._build_test_env(test_config) @@ -2260,15 +2658,20 @@ def test_max_requests(self, max_requests: int | None): context = env.engine.context if max_requests is None: assert context.max_requests == 816 - assert step_count == 23 else: assert max_requests < len(env.requests), ( f"Test is only useful if max_requests ({max_requests}) < " f"num_requests ({len(env.requests)})." ) assert context.max_requests == 4 - assert step_count == 35 - assert context.kv_block_allocator.active_count == 655 + # Exact step counts and KV occupancy depend on sampled token sequences. + # With DP-offset sampling seeds, only DP rank 0 matches the golden seed. + if parallel_state.get_data_parallel_rank() == 0: + if max_requests is None: + assert step_count == 23 + else: + assert step_count == 35 + assert context.kv_block_allocator.active_count == 655 @pytest.mark.internal @pytest.mark.skipif( @@ -4057,8 +4460,8 @@ def test_speculative_decoding_non_greedy_with_top_n_logprobs(self): env.engine.controller.tokenizer.detokenize = lambda tokens, **kw: f"tok_{tokens[0]}" - # top_n must be >= top_k so the sampled token is guaranteed to appear - # in the top-n dict for the consistency check below. + # top_n must be >= top_k so the top_k-sampled token is guaranteed to + # have a higher probability than the least probable token in top_n. top_n = 10 num_requests = 3 prompt_lengths = [4, 6, 8] @@ -4097,20 +4500,30 @@ def test_speculative_decoding_non_greedy_with_top_n_logprobs(self): assert isinstance(top_n_dict, dict) assert 0 < len(top_n_dict) <= top_n - # Consistency: selected token's log prob should appear in top-n. + # Consistency: selected token's log prob should appear in top-n when the + # sampled token is among the strict top-n indices. With nearly-tied logits + # (random models), top-k filtering keeps all ties at the cutoff, so a sampled + # token can fall outside a separate top-n index list. In that case the + # selected logprob must still be no worse than the weakest top-n entry. if req.generated_log_probs is not None: for j, (lp, top_n_dict, token_id) in enumerate( zip(req.generated_log_probs, req.generated_top_n_logprobs, req.generated_tokens) ): token_str = env.engine.controller.tokenizer.detokenize([token_id]) - assert token_str in top_n_dict, ( - f"Request {req.request_id}, token {j}: " - f"selected token '{token_str}' not in top-n" - ) - assert abs(lp - top_n_dict[token_str]) < 0.01, ( - f"Request {req.request_id}, token {j}: " - f"log_prob {lp} vs top-n {top_n_dict[token_str]}" - ) + if token_str in top_n_dict: + # Sampled token is in Top N. + assert abs(lp - top_n_dict[token_str]) < 0.01, ( + f"Request {req.request_id}, token {j}: " + f"log_prob {lp} vs top-n {top_n_dict[token_str]}" + ) + else: + # Sampled token is not in the Top N. It must be a tie. + # Check that it is at least as probable as Top N tokens. + assert lp + 0.01 >= min(top_n_dict.values()), ( + f"Request {req.request_id}, token {j}: " + f"selected token '{token_str}' log_prob {lp} is worse than " + f"top-n minimum {min(top_n_dict.values())}" + ) @pytest.mark.internal @pytest.mark.skipif( @@ -4214,7 +4627,8 @@ def mock_compute_mtp_wrong( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" ) @torch.inference_mode() - def test_speculative_decoding_logprobs_with_stop_word_trim(self): + @pytest.mark.parametrize("async_sched_mode", list(AsyncScheduleMode)) + def test_speculative_decoding_logprobs_with_stop_word_trim(self, async_sched_mode): """Test that log probs are correctly trimmed when a stop word lands in the middle of a speculative batch. @@ -4223,6 +4637,9 @@ def test_speculative_decoding_logprobs_with_stop_word_trim(self): generates [5, 6, 7] in one step, token 7 is truncated. The corresponding log prob for token 7 must also be removed so that len(generated_log_probs) == len(generated_tokens). + + Args: + async_sched_mode (AsyncScheduleMode): Scheduling mode under test. """ test_config = DynamicEngineTestConfig( num_requests=0, @@ -4232,6 +4649,7 @@ def test_speculative_decoding_logprobs_with_stop_word_trim(self): num_speculative_tokens=2, materialize_only_last_token_logits=False, model_provider="gpt", + async_sched_mode=async_sched_mode, ) env = self._build_test_env(test_config) @@ -4274,6 +4692,7 @@ def mock_compute_mtp_single_step( detokenize_stop_sequence=True, return_log_probs=True, top_k=1, + top_n_logprobs=2, ), ) @@ -4288,6 +4707,7 @@ def mock_compute_mtp_single_step( finished_req = finished_records[0].merge() assert finished_req.status == Status.COMPLETED + assert finished_req.generated_tokens == [5, 6] assert finished_req.generated_tokens[-1] == 6, ( f"Expected last token to be stop word 6, " f"got {finished_req.generated_tokens[-1]}. " @@ -4307,6 +4727,9 @@ def mock_compute_mtp_single_step( assert isinstance(lp, float) assert lp <= 0.0, f"Token {j}: log prob {lp} > 0" + assert finished_req.generated_top_n_logprobs is not None + assert len(finished_req.generated_top_n_logprobs) == len(finished_req.generated_tokens) + @pytest.mark.internal @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" @@ -4853,6 +5276,49 @@ def _build_test_env(cls, test_config): ) return super()._build_test_env(test_config) + @pytest.mark.internal + @pytest.mark.skipif( + not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" + ) + @pytest.mark.parametrize( + "async_sched_mode", [AsyncScheduleMode.LEGACY, AsyncScheduleMode.ASYNC] + ) + @torch.inference_mode() + def test_non_greedy_sampling_with_mamba_mtp_ep(self, async_sched_mode): + """Run cumulative Mamba, MTP, and EP sampling support to completion. + + Args: + async_sched_mode (AsyncScheduleMode): Scheduling mode under test. + """ + skip_if_mamba_sequence_packing_not_available("hybrid") + if int(os.environ.get("WORLD_SIZE", "1")) < 2: + pytest.skip("Test requires at least 2 GPUs") + + env = self._run_test( + num_requests=2, + min_prompt_length=4, + max_prompt_length=4, + num_tokens_to_generate=4, + num_gap_steps=0, + use_fixed_output_lengths=True, + model_provider="hybrid", + expert_model_parallel_size=2, + num_speculative_tokens=1, + sampling_backend="torch", + temperature=0.8, + top_k=8, + return_log_probs=True, + skip_prompt_log_probs=True, + context_max_requests=8, + async_sched_mode=async_sched_mode, + ) + + assert all(request.status == Status.COMPLETED for request in env.requests) + assert all( + len(request.generated_log_probs) == len(request.generated_tokens) + for request in env.requests + ) + @pytest.mark.internal @pytest.mark.skipif( not is_fa_min_version("2.7.3"), reason="need latest flash attn for dynamic batching" diff --git a/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py b/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py index f1ad8a84517..b8eb78b2403 100644 --- a/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py +++ b/tests/unit_tests/inference/engines/test_dynamic_engine_async_sched.py @@ -1,5 +1,7 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +import asyncio +from collections import deque from types import SimpleNamespace from unittest import mock @@ -7,29 +9,43 @@ from megatron.core.inference.config import AsyncScheduleMode from megatron.core.inference.engines import DynamicInferenceEngine +from megatron.core.inference.engines.dynamic_engine import EngineState, _get_decode_only_log_state from megatron.core.inference.sampling_params import SamplingParams +from megatron.core.inference.text_generation_controllers.text_generation_controller import ( + DecodeOnly, + DynamicBatchControllerStepResult, +) -def _make_engine(async_sched_mode=AsyncScheduleMode.SERIAL, **overrides): +def _make_engine(async_sched_mode=AsyncScheduleMode.ASYNC, **overrides): engine = DynamicInferenceEngine.__new__(DynamicInferenceEngine) context = SimpleNamespace( config=SimpleNamespace(async_sched_mode=async_sched_mode), is_hybrid_model=False, enable_prefix_caching=False, + num_prefill_requests=0, + can_prepare_requests=mock.Mock(return_value=True), + active_token_count=0, + max_tokens=8, + chunked_prefill_request_id=-1, ) model_config = SimpleNamespace( expert_model_parallel_size=1, num_moe_experts=None, moe_enable_routing_replay=False ) engine.context = context engine.controller = SimpleNamespace( - inference_wrapped_model=SimpleNamespace(model=SimpleNamespace(config=model_config)) + inference_wrapped_model=SimpleNamespace(model=SimpleNamespace(config=model_config)), + num_mtp_depths=0, ) + engine.enable_chunked_prefill = False engine.num_speculative_tokens = 0 engine.materialize_only_last_token_logits = True for name, value in overrides.items(): if name.startswith("context_"): setattr(context, name.removeprefix("context_"), value) + elif name.startswith("controller_"): + setattr(engine.controller, name.removeprefix("controller_"), value) elif name.startswith("model_config_"): setattr(model_config, name.removeprefix("model_config_"), value) else: @@ -42,12 +58,32 @@ def _make_engine(async_sched_mode=AsyncScheduleMode.SERIAL, **overrides): [ ({"async_sched_mode": AsyncScheduleMode.LEGACY, "num_speculative_tokens": 1}, False), ({}, False), + ({"enable_chunked_prefill": True}, False), ({"num_speculative_tokens": 1}, True), - ({"context_is_hybrid_model": True}, True), - ({"context_enable_prefix_caching": True}, True), - ({"materialize_only_last_token_logits": False}, True), - ({"model_config_expert_model_parallel_size": 2}, True), - ({"model_config_num_moe_experts": 4}, True), + ({"num_speculative_tokens": 1, "controller_num_mtp_depths": 1}, False), + ({"context_is_hybrid_model": True}, False), + ( + { + "context_is_hybrid_model": True, + "num_speculative_tokens": 1, + "controller_num_mtp_depths": 1, + "model_config_expert_model_parallel_size": 2, + "model_config_num_moe_experts": 4, + }, + False, + ), + ({"context_enable_prefix_caching": True}, False), + ( + { + "enable_chunked_prefill": True, + "context_enable_prefix_caching": True, + "context_is_hybrid_model": True, + }, + False, + ), + ({"materialize_only_last_token_logits": False}, False), + ({"model_config_expert_model_parallel_size": 2}, False), + ({"model_config_num_moe_experts": 4}, False), ({"model_config_moe_enable_routing_replay": True}, True), ], ) @@ -63,38 +99,266 @@ def test_validate_async_sched_support_for_config(overrides, should_raise): @pytest.mark.parametrize( - "async_sched_mode, sampling_params, should_raise", + "can_prepare, has_waiting, availability, expected", [ - (AsyncScheduleMode.LEGACY, SamplingParams(top_k=0, top_p=0.5), False), - (AsyncScheduleMode.SERIAL, SamplingParams(top_k=1, top_p=0.0), False), - (AsyncScheduleMode.SERIAL, SamplingParams(top_k=0, top_p=0.0), True), - (AsyncScheduleMode.SERIAL, SamplingParams(top_k=1, top_p=0.5), True), - (AsyncScheduleMode.SERIAL, SamplingParams(top_k=1, top_p=0.0, return_log_probs=True), True), - (AsyncScheduleMode.SERIAL, SamplingParams(top_k=1, top_p=0.0, top_n_logprobs=1), True), - (AsyncScheduleMode.SERIAL, SamplingParams(top_k=1, top_p=0.0, stop_words=["END"]), True), + (False, False, (False, False, False), False), + (True, False, (True, True, True), True), + (True, True, (False, True, True), True), + (True, True, (True, True, True), False), ], ) -def test_validate_async_sched_support_for_request(async_sched_mode, sampling_params, should_raise): - """Ensure engine request validation accepts only supported async scheduling requests.""" - engine = _make_engine(async_sched_mode=async_sched_mode) - request = SimpleNamespace(sampling_params=sampling_params) +def test_should_run_async_sched_overlap(can_prepare, has_waiting, availability, expected): + """The overlap probe observes prefill eligibility without admitting the request.""" + engine = _make_engine() + engine.context.can_prepare_requests.return_value = can_prepare + engine.context.check_availability = mock.Mock(return_value=availability) + engine.waiting_request_ids = deque([10] if has_waiting else []) + request = SimpleNamespace(remaining_prompt_tokens=[1, 2], cg_wait_iters=3) + engine.get_request = mock.Mock(return_value=request) + engine._cg_admission_gating_active = mock.Mock(return_value=False) - if should_raise: - with pytest.raises(ValueError, match="Async scheduling"): - engine._validate_async_sched_support_for_request(request) + assert engine._should_run_async_sched_overlap() is expected + engine.context.can_prepare_requests.assert_called_once_with() + assert list(engine.waiting_request_ids) == ([10] if has_waiting else []) + assert request.cg_wait_iters == 3 + + +def test_async_sched_overlap_probe_uses_non_mutating_cuda_graph_match(): + """A scheduling probe does not update CUDA-graph wait accounting.""" + engine = _make_engine() + engine.context.active_token_count = 2 + engine.context.num_prefill_requests = 0 + engine.context.num_decode_requests = 2 + engine.context.check_availability = mock.Mock(return_value=(True, True, True)) + engine._cg_admission_gating_active = mock.Mock(return_value=True) + engine._matches_cg_admission = mock.Mock(return_value=False) + engine._cg_admission_check = mock.Mock() + request = SimpleNamespace(remaining_prompt_tokens=[1, 2], cg_wait_iters=7) + + assert not engine._can_schedule_non_chunked_prefill(request, record_cg_wait=False) + engine._matches_cg_admission.assert_called_once() + engine._cg_admission_check.assert_not_called() + assert request.cg_wait_iters == 7 + + +@pytest.mark.parametrize( + "availability, active_token_count, chunked_prefill_request_id, expected", + [ + ((True, False, True), 7, -1, True), + ((False, False, True), 7, 10, True), + ((True, True, False), 7, -1, False), + ((True, True, True), 8, -1, False), + ], +) +def test_can_schedule_chunked_prefill( + availability, active_token_count, chunked_prefill_request_id, expected +): + """The chunk probe requires request, KV-cache, and partial-token capacity.""" + engine = _make_engine(enable_chunked_prefill=True) + engine.context.active_token_count = active_token_count + engine.context.chunked_prefill_request_id = chunked_prefill_request_id + engine.context.check_availability = mock.Mock(return_value=availability) + request = SimpleNamespace(request_id=10) + + assert engine._can_schedule_chunked_prefill(request) is expected + + +def test_async_sched_overlap_probe_routes_schedulable_chunk_to_no_overlap(): + """A schedulable chunk is admitted only after no-overlap lifecycle bookkeeping.""" + engine = _make_engine(enable_chunked_prefill=True) + engine.context.active_token_count = 2 + engine.context.check_availability = mock.Mock(return_value=(True, False, True)) + engine.waiting_request_ids = deque([10]) + engine.get_request = mock.Mock(return_value=SimpleNamespace(request_id=10)) + + assert not engine._should_run_async_sched_overlap() + + +@pytest.mark.parametrize( + "mode, run_async_overlap, decode_only, primer_only, expected_schedule_calls, " + "expected_nvtx_range", + [ + ( + AsyncScheduleMode.LEGACY, + None, + DecodeOnly(consumed=False, launched=False), + False, + 1, + "Prefill", + ), + ( + AsyncScheduleMode.LEGACY, + None, + DecodeOnly(consumed=True, launched=True), + False, + 1, + "Decode", + ), + ( + AsyncScheduleMode.ASYNC, + True, + DecodeOnly(consumed=True, launched=True), + False, + 0, + "AsyncOverlap", + ), + ( + AsyncScheduleMode.ASYNC, + False, + DecodeOnly(consumed=False, launched=True), + False, + 0, + "AsyncNoOverlap", + ), + ( + AsyncScheduleMode.ASYNC, + False, + DecodeOnly(consumed=None, launched=False), + True, + 0, + "AsyncNoOverlap", + ), + ( + AsyncScheduleMode.ASYNC, + False, + DecodeOnly(consumed=True, launched=None), + False, + 0, + "AsyncNoOverlap", + ), + ], +) +def test_async_forward_routes_one_controller_iteration( + mode, run_async_overlap, decode_only, primer_only, expected_schedule_calls, expected_nvtx_range +): + """Primer-only work crosses the engine boundary without an internal controller loop.""" + engine = DynamicInferenceEngine.__new__(DynamicInferenceEngine) + engine.state = EngineState.RUNNING + engine.logging_step_interval = 0 + engine.metrics_writer = None + engine.schedule_waiting_requests = mock.Mock() + engine._should_run_async_sched_overlap = mock.Mock(return_value=run_async_overlap) + engine.context = SimpleNamespace( + config=SimpleNamespace(async_sched_mode=mode), + step_count=4, + prefix_cache_lru_clock=7, + active_token_count=2, + num_prefill_requests=1 if expected_nvtx_range == "Prefill" else 0, + chunked_prefill_request_id=17, + is_decode_only=mock.Mock(return_value=decode_only.launched), + ) + output = None if primer_only else {"sample": "tokens"} + engine.controller = SimpleNamespace( + async_generate_output_tokens_dynamic_batch=mock.AsyncMock( + return_value=DynamicBatchControllerStepResult( + decode_only=decode_only, output=output, primer_only=primer_only + ) + ) + ) + + with ( + mock.patch( + "megatron.core.inference.engines.dynamic_engine.nvtx_range_push" + ) as nvtx_range_push, + mock.patch( + "megatron.core.inference.engines.dynamic_engine.nvtx_range_pop" + ) as nvtx_range_pop, + ): + result, context_state, _ = asyncio.run(engine.async_forward()) + + assert result is output + assert context_state["decode_only"] == decode_only + assert context_state["chunked_prefill_request_id"] == 17 + assert engine.decode_only == decode_only + assert not hasattr(engine, "is_decode_only") + assert engine.context.step_count == 5 + assert engine.context.prefix_cache_lru_clock == 8 + assert engine.schedule_waiting_requests.call_count == expected_schedule_calls + nvtx_range_push.assert_called_once_with(expected_nvtx_range) + nvtx_range_pop.assert_called_once_with(expected_nvtx_range) + if mode == AsyncScheduleMode.LEGACY: + engine._should_run_async_sched_overlap.assert_not_called() + engine.controller.async_generate_output_tokens_dynamic_batch.assert_awaited_once_with() else: - engine._validate_async_sched_support_for_request(request) + engine._should_run_async_sched_overlap.assert_called_once_with() + engine.controller.async_generate_output_tokens_dynamic_batch.assert_awaited_once_with( + run_async_overlap=run_async_overlap, + schedule_waiting_requests=( + None if run_async_overlap else engine.schedule_waiting_requests + ), + ) + engine.context.is_decode_only.assert_not_called() -def test_add_request_runs_async_sched_request_validation(): - """Ensure request validation is called before mutating engine request state.""" +def test_async_bookkeep_uses_consumed_chunked_prefill_request_id(): + """Post-processing classifies output using the chunk ID from its consumed forward.""" engine = DynamicInferenceEngine.__new__(DynamicInferenceEngine) - engine._validate_async_sched_support_for_request = mock.Mock( - side_effect=RuntimeError("validated") + engine.track_paused_request_events = False + engine.post_process_requests = mock.Mock(return_value=([10], [])) + engine.failed_request_ids = set() + engine.requests = {} + engine.use_coordinator = False + engine.context = SimpleNamespace(enable_prefix_caching=False, step_count=1) + engine.logging_step_interval = 0 + engine.num_speculative_tokens = 0 + step_result = { + "active_request_ids": [10], + "finished_request_ids": [], + "sample": [20], + "accepted_tokens": None, + "log_probs": None, + "cuda_graph_request_count": None, + } + context_state = { + "active_token_count": 4, + "step_count": 0, + "chunked_prefill_request_id": 10, + "kv_stats": None, + } + + with ( + mock.patch("megatron.core.inference.engines.dynamic_engine.nvtx_range_push"), + mock.patch("megatron.core.inference.engines.dynamic_engine.nvtx_range_pop"), + ): + asyncio.run(engine.async_bookkeep(step_result, context_state, 0.0)) + + assert ( + engine.post_process_requests.call_args.kwargs["consumed_chunked_prefill_request_id"] == 10 ) - request = SimpleNamespace(request_id=10) - with pytest.raises(RuntimeError, match="validated"): - engine._add_request(request) - engine._validate_async_sched_support_for_request.assert_called_once_with(request) +@pytest.mark.parametrize( + "mode, decode_only, expected", + [ + ( + AsyncScheduleMode.LEGACY, + DecodeOnly(consumed=False, launched=False), + ("non-decode", False), + ), + (AsyncScheduleMode.LEGACY, DecodeOnly(consumed=True, launched=True), ("decode", True)), + ( + AsyncScheduleMode.ASYNC, + DecodeOnly(consumed=False, launched=False), + ("non-decode", False), + ), + (AsyncScheduleMode.ASYNC, DecodeOnly(consumed=True, launched=True), ("decode", True)), + ( + AsyncScheduleMode.ASYNC, + DecodeOnly(consumed=False, launched=True), + ("decode (prev: non-decode)", True), + ), + ( + AsyncScheduleMode.ASYNC, + DecodeOnly(consumed=True, launched=False), + ("non-decode (prev: decode)", False), + ), + (AsyncScheduleMode.ASYNC, DecodeOnly(consumed=None, launched=False), ("non-decode", False)), + (AsyncScheduleMode.ASYNC, DecodeOnly(consumed=None, launched=True), ("decode", True)), + (AsyncScheduleMode.ASYNC, DecodeOnly(consumed=False, launched=None), ("non-decode", False)), + (AsyncScheduleMode.ASYNC, DecodeOnly(consumed=True, launched=None), ("decode", True)), + (AsyncScheduleMode.ASYNC, DecodeOnly(consumed=None, launched=None), ("idle", None)), + ], +) +def test_get_decode_only_log_state(mode, decode_only, expected): + """Console logging reports transitions and colors the latest available phase.""" + assert _get_decode_only_log_state(mode, decode_only) == expected diff --git a/tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py b/tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py index 92890b22b53..007f4a2c2b7 100644 --- a/tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py +++ b/tests/unit_tests/inference/engines/test_hybrid_prefix_caching_e2e.py @@ -40,6 +40,7 @@ from megatron.core import parallel_state from megatron.core.inference.config import ( + AsyncScheduleMode, InferenceConfig, MambaInferenceStateConfig, PrefixCachingEvictionPolicy, @@ -189,6 +190,10 @@ def _build_engine( prefix_caching_mamba_gb=0.05, request_rounder=4, num_cuda_graphs=None, + enable_chunked_prefill=False, + max_tokens=None, + max_requests=None, + async_sched_mode=AsyncScheduleMode.LEGACY, ): set_rounder(request_rounder) inference_config_kwargs = dict( @@ -196,12 +201,18 @@ def _build_engine( buffer_size_gb=buffer_size_gb, block_size_tokens=BLOCK_SIZE, mamba_inference_state_config=mamba_config, - materialize_only_last_token_logits=False, + materialize_only_last_token_logits=async_sched_mode == AsyncScheduleMode.ASYNC, enable_prefix_caching=enable_prefix_caching, + enable_chunked_prefill=enable_chunked_prefill, unified_memory_level=0, num_cuda_graphs=num_cuda_graphs, sampling_backend='torch', + async_sched_mode=async_sched_mode, ) + if max_tokens is not None: + inference_config_kwargs["max_tokens"] = max_tokens + if max_requests is not None: + inference_config_kwargs["max_requests"] = max_requests if enable_prefix_caching: inference_config_kwargs.update( prefix_caching_eviction_policy=PrefixCachingEvictionPolicy.LRU, @@ -425,6 +436,31 @@ def test_mamba_prefix_caching_e2e(self): ), f"req {req_id}: pc=off {off_outputs[req_id]} != pc=on {on_outputs[req_id]}" assert off_prefill == 3800 and on_prefill == 2008 and on_prefill < off_prefill + @torch.inference_mode() + def test_async_sched_mamba_prefix_caching_with_chunked_prefill_e2e(self): + """Async combined chunking and Mamba prefix caching matches legacy output.""" + skip_if_mamba_sequence_packing_not_available() + model = self._create_model() + mamba_config = MambaInferenceStateConfig.from_model(model) + prompts = self._create_prompts()[:3] + + legacy_outputs, legacy_prefill = self._run_simple( + model, mamba_config, prompts, enable_pc=False + ) + async_outputs, async_prefill = self._run_simple( + model, + mamba_config, + prompts, + enable_pc=True, + enable_chunked_prefill=True, + max_tokens=400, + max_requests=4, + async_sched_mode=AsyncScheduleMode.ASYNC, + ) + + assert async_outputs == legacy_outputs + assert async_prefill < legacy_prefill + @pytest.mark.parametrize("num_cuda_graphs", [None, 2]) @torch.inference_mode() def test_mamba_prefix_caching_multi_group_e2e(self, num_cuda_graphs): @@ -619,3 +655,75 @@ def _run_one(req_id, prompt): assert req_G._mamba_num_matched_blocks == 0 assert h_E0 in ctx.mamba_slot_allocator.hash_to_block_id assert finished[0] == finished[2] + + @torch.inference_mode() + def test_mamba_chunked_prefill_unaligned_boundary_snapshot(self): + """Chunked prefill snapshots Mamba state at the last block boundary. + + ``compute_and_store_offsets`` records a Mamba state snapshot at a KV-block + boundary only when that boundary is a whole multiple of the SSM chunk size + measured from the start of the current prefill chunk. Because the chunk + start equals ``finished_chunk_token_count`` on continuation chunks, this + holds exactly when every chunk boundary is block-aligned. + + Here ``max_tokens`` (300) is intentionally not a multiple of the block size + (256), so the request spans several chunks and its last full-block boundary + (token 768) lands in a continuation chunk. The scheduler keeps each chunk + boundary block-aligned, so the final chunk begins at token 512 and the + token-768 snapshot is extracted and committed. A second request sharing the + 768-token prefix then restores that state and skips those blocks. + """ + skip_if_mamba_sequence_packing_not_available() + model = self._create_model() + mamba_config = MambaInferenceStateConfig.from_model(model) + + device = torch.cuda.current_device() + # 800-token prompt -> 3 full blocks (256/512/768) + a 32-token tail. + # The last full-block boundary (768) falls in the final continuation chunk. + prompt = torch.arange(9000, 9800, dtype=torch.int64, device=device) + assert len(prompt) == 800 + + engine = self._build_engine( + model, + mamba_config, + enable_prefix_caching=True, + enable_chunked_prefill=True, + max_tokens=300, # not a multiple of BLOCK_SIZE (256) -> forces unaligned cuts + max_requests=4, + request_rounder=4, + ) + ctx = engine.context + # Sanity: the prompt genuinely spans multiple prefill chunks. + assert ctx.max_tokens < len(prompt) + + # --- Seed request: fills the cache, no prior matches. --- + seed = self._make_request(0, prompt, enable_pc=True, num_tokens=4) + engine._add_request(seed) + while engine.has_unfinished_requests(): + engine.step_modern() + # Seed has no prior cache, so no Mamba blocks are matched during its prefill. + assert seed._mamba_num_matched_blocks == 0 + + # block index 2 == the boundary at token 768 (768 // 256 - 1). The final + # chunk begins block-aligned at token 512, so (768 - 512) % 128 == 0 and + # the state at this boundary is extracted and committed. + assert len(seed.precomputed_block_hashes) == 3 + last_block_hash = seed.precomputed_block_hashes[2] + assert ( + last_block_hash in ctx.mamba_slot_allocator.hash_to_block_id + ), "Mamba snapshot at the last block boundary (token 768) was not recorded." + + # --- Reuse request: shares the full 768-token prefix, should restore the + # cached Mamba state and skip those blocks entirely. --- + reuse_prompt = torch.cat( + [prompt[:768], torch.arange(9800, 9900, dtype=torch.int64, device=device)] + ) + reuse = self._make_request(1, reuse_prompt, enable_pc=True, num_tokens=4) + engine._add_request(reuse) + while engine.has_unfinished_requests(): + engine.step_modern() + + assert reuse._mamba_num_matched_blocks == 3, ( + "Reuse request should restore Mamba state from the token-768 snapshot " + f"(3 matched blocks), got {reuse._mamba_num_matched_blocks}." + ) diff --git a/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py b/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py index c89209699b3..1cc0adf4d1b 100644 --- a/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py +++ b/tests/unit_tests/inference/engines/test_prefix_caching_cuda_graphs.py @@ -12,7 +12,6 @@ 2. context.using_cuda_graph_this_step() returned True at expected steps. """ -import os import random import types @@ -44,7 +43,7 @@ from megatron.core.transformer.cuda_graphs import CudaGraphManager, _CudagraphGlobalRecord from megatron.core.transformer.transformer_config import TransformerConfig from megatron.core.utils import is_fa_min_version -from tests.unit_tests.test_utilities import Utils +from tests.unit_tests.test_utilities import Utils, clear_nvte_env_vars BLOCK_SIZE = 256 VOCAB_SIZE = 10000 @@ -260,6 +259,7 @@ def _step_and_log(): return finished, step_log + @pytest.mark.flaky_in_dev # Issue #6130 @pytest.mark.parametrize("model_type", ["transformer", "hybrid"]) @pytest.mark.parametrize("batch_structure", ["prefill", "decode", "mixed"]) @torch.inference_mode() @@ -319,6 +319,11 @@ class TestHybridChunkedPrefillIntermediateState: @classmethod def setup_class(cls): Utils.initialize_model_parallel() + random.seed(123) + torch.manual_seed(123) + model_parallel_cuda_manual_seed( + seed=123, inference_rng_tracker=True, use_cudagraphable_rng=False, force_reset_rng=True + ) @classmethod def teardown_class(cls): @@ -443,16 +448,7 @@ def test_hybrid_chunked_prefill_intermediate_state(self): if not sequence_packing_available: pytest.skip(reason) - # Clear NVTE env vars set by conftest set_env fixture. - os.environ.pop('NVTE_FLASH_ATTN', None) - os.environ.pop('NVTE_FUSED_ATTN', None) - os.environ.pop('NVTE_UNFUSED_ATTN', None) - - random.seed(123) - torch.manual_seed(123) - model_parallel_cuda_manual_seed( - seed=123, inference_rng_tracker=True, use_cudagraphable_rng=False, force_reset_rng=True - ) + clear_nvte_env_vars() # conftest's set_env fixture re-sets these per test model = self._create_hybrid_model() mamba_config = MambaInferenceStateConfig.from_model(model) @@ -534,3 +530,50 @@ def collect_finished(result): f"req {req_id}: baseline {baseline_outputs[req_id]} != " f"test {test_outputs[req_id]}" ) + + @torch.inference_mode() + def test_prefill_shorter_than_conv_window(self): + """A prefill captured into a CUDA graph whose token bucket is smaller than the + Mamba conv window (d_conv) generates correctly. + + Conv-state extraction gathers d_conv positions per slot, and unused slots use + abs_position == d_conv (gather indices up to d_conv-1). The CUDA-graph bucket + list always includes a size-1 (tp_size) graph, so a prompt shorter than d_conv + is captured at a bucket whose token layout is shorter than the gather window. + CUDA graphs (num_cuda_graphs) are required to exercise this capture path. + """ + sequence_packing_available, reason = _check_mamba_sequence_packing_support() + if not sequence_packing_available: + pytest.skip(reason) + + clear_nvte_env_vars() # conftest's set_env fixture re-sets these per test + + model = self._create_hybrid_model(num_cuda_graphs=2) + mamba_config = MambaInferenceStateConfig.from_model(model) + device = torch.cuda.current_device() + + d_conv = mamba_config.conv_states_shape[-1] + if d_conv < 2: + pytest.skip(f"d_conv={d_conv} too small to exercise a sub-window prefill") + + # Prompt shorter than the conv window: its prefill chunk snaps to a CUDA-graph + # bucket < d_conv, so the captured graph's token layout is < d_conv. + engine = self._build_engine( + model, + mamba_config, + enable_prefix_caching=True, + enable_chunked_prefill=True, + num_cuda_graphs=2, + ) + short_prompt = torch.arange(0, d_conv - 1, dtype=torch.int64, device=device) + + engine._add_request(self._make_request(0, short_prompt, enable_pc=True)) + outputs = {} + while engine.has_unfinished_requests(): + result = engine.step_modern() + for record in result["finished_request_records"]: + merged = record.merge() + outputs[merged.request_id] = list(merged.generated_tokens) + + # Generation completes and produces the requested number of tokens. + assert len(outputs[0]) == NUM_TOKENS_TO_GENERATE diff --git a/tests/unit_tests/inference/engines/test_static_engine.py b/tests/unit_tests/inference/engines/test_static_engine.py index e9befc290fe..d40766b6a26 100644 --- a/tests/unit_tests/inference/engines/test_static_engine.py +++ b/tests/unit_tests/inference/engines/test_static_engine.py @@ -212,6 +212,10 @@ def test_generate_dynamic(self, batch_size: int, num_trials: int, empty_prompt: async def test_streaming(self): self.setup_engine(legacy=True) + # Possible for a rank to not generate any tokens, i.e. EOD only. + # Make that impossible when testing streaming. + self.mock_tokenizer.eod = self.vocab_size + async def collect_stream(stream_generator, num_tokens_to_generate): prev_log_probs = None prev_text = "" diff --git a/tests/unit_tests/inference/high_level_api/test_apis.py b/tests/unit_tests/inference/high_level_api/test_apis.py index c9fbad3de2c..ff6721a4f0d 100644 --- a/tests/unit_tests/inference/high_level_api/test_apis.py +++ b/tests/unit_tests/inference/high_level_api/test_apis.py @@ -13,6 +13,7 @@ from megatron.core.inference.apis._llm_base import _MegatronLLMBase from megatron.core.inference.apis.async_llm import MegatronAsyncLLM from megatron.core.inference.apis.llm import MegatronLLM +from megatron.core.inference.apis.serve_config import ServeConfig @pytest.fixture @@ -50,8 +51,7 @@ def _make_worker_instance(cls): obj._loop_manager = None obj._coord_runtime = None obj._shutdown_called = False - if cls is MegatronAsyncLLM: - obj._serve_started = False + obj._serve_started = False return obj @@ -130,6 +130,51 @@ async def test_async_generate_raises_on_worker_rank(self): with pytest.raises(RuntimeError, match="primary rank"): await llm.generate("hello") + def test_bridge_and_serve_raise_in_direct_mode(self, mock_pipeline, fake_model_and_tokenizer): + model, tok = fake_model_and_tokenizer + llm = MegatronLLM(model=model, tokenizer=tok, use_coordinator=False) + with pytest.raises(ValueError, match="use_coordinator=True"): + llm.serve(ServeConfig()) + + async def coro(): + return 1 # pragma: no cover + + for method in (llm.run_sync, llm.submit): + c = coro() + with pytest.raises(RuntimeError, match="use_coordinator=True"): + method(c) + c.close() + + def test_sync_serve_nonblocking_worker_rank_noops(self): + """Worker ranks skip the HTTP setup; ``blocking=False`` returns + immediately without touching the runtime.""" + llm = _make_worker_instance(MegatronLLM) + llm.serve(ServeConfig(), blocking=False) + assert llm._serve_started is False + + def test_sync_serve_primary_rank_starts_frontend(self, monkeypatch): + """Primary rank starts the HTTP frontend against the coordinator + address and records ``_serve_started`` for shutdown teardown.""" + tgs = pytest.importorskip( + "megatron.core.inference.text_generation_server.dynamic_text_gen_server" + ".text_generation_server" + ) + import torch.distributed as dist + + llm = _make_worker_instance(MegatronLLM) + llm._is_primary_rank = True + llm._coord_runtime = MagicMock() + llm._coord_runtime.coord_addr = "tcp://coord:5555" + + started = {} + monkeypatch.setattr(dist, "get_rank", lambda: 0) + monkeypatch.setattr(tgs, "start_text_gen_server", lambda **kw: started.update(kw)) + + llm.serve(ServeConfig(port=1234), blocking=False) + assert llm._serve_started is True + assert started["coordinator_addr"] == "tcp://coord:5555" + assert started["server_port"] == 1234 + class TestNormalizePrompts: """Input-shape normalization (str / list[int] / list[str] / list[list[int]]).""" diff --git a/tests/unit_tests/inference/high_level_api/test_event_loop_manager.py b/tests/unit_tests/inference/high_level_api/test_event_loop_manager.py index 9d647dddad0..2b903f5c082 100644 --- a/tests/unit_tests/inference/high_level_api/test_event_loop_manager.py +++ b/tests/unit_tests/inference/high_level_api/test_event_loop_manager.py @@ -1,5 +1,6 @@ # Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +import asyncio import threading import pytest @@ -76,3 +77,23 @@ async def deadlock_attempt(): mgr.submit(deadlock_attempt()).result() finally: mgr.stop() + + def test_run_sync_from_foreign_running_loop_returns_result(self): + """A thread whose own event loop is running (e.g. an embedder's + dispatch loop) may call ``run_sync``: the caller's loop stalls while + the coroutine executes on the background loop and the result is + handed back.""" + mgr = _EventLoopManager() + mgr.start() + try: + + async def inner(): + return 42 + + async def foreign_caller(): + # Runs on a fresh caller-owned loop, not mgr._loop. + return mgr.run_sync(inner()) + + assert asyncio.run(foreign_caller()) == 42 + finally: + mgr.stop() diff --git a/tests/unit_tests/inference/test_async_sched_output_metrics.py b/tests/unit_tests/inference/test_async_sched_output_metrics.py index 5c327cfb7f6..0c12659286b 100644 --- a/tests/unit_tests/inference/test_async_sched_output_metrics.py +++ b/tests/unit_tests/inference/test_async_sched_output_metrics.py @@ -3,7 +3,11 @@ import json from argparse import Namespace from types import SimpleNamespace +from unittest import mock +import pytest + +from examples.inference import utils as inference_utils from examples.inference.offline_inference import _capture_engine_stats from examples.inference.utils import dump_inference_results_to_json from tests.functional_tests.python_test_utils.test_inference_regular_pipeline import ( @@ -53,6 +57,25 @@ def test_inference_comparator_ignores_async_sched_counters(): assert "async_sched_compaction_step_count" in _NON_REQUEST_TOP_LEVEL_KEYS +@pytest.mark.parametrize(("prompt_tokens", "prompt_length"), [([1, 2, 3], None), (None, 3)]) +def test_print_unique_prompts_and_outputs_uses_available_prompt_length( + capsys, prompt_tokens, prompt_length +): + """Ensure reporting supports direct and coordinator inference results.""" + request = SimpleNamespace( + prompt="prompt", + prompt_tokens=prompt_tokens, + prompt_length=prompt_length, + generated_text="generated", + generated_tokens=[4], + events=[], + ) + + inference_utils.print_unique_prompts_and_outputs([request]) + + assert "[n 1, l 3] prompt" in capsys.readouterr().out + + def test_capture_engine_stats_includes_async_sched_counters(): """Ensure offline reporting captures async scheduling counters from the engine context.""" context = SimpleNamespace( @@ -70,3 +93,41 @@ def test_capture_engine_stats_includes_async_sched_counters(): "async_sched_compaction_step_count": 4, "capture_stats": {"graphs": 5}, } + + +@pytest.mark.parametrize( + ("do_broadcast", "distributed_initialized", "world_size"), + [(False, True, 2), (True, False, 2), (True, True, 1)], +) +def test_get_curr_time_avoids_cuda_when_rank_sync_is_unnecessary( + monkeypatch, do_broadcast, distributed_initialized, world_size +): + """Ensure local timing never synchronizes the CUDA compute stream.""" + monkeypatch.setattr(inference_utils.time, "time_ns", lambda: 123_000_000_000) + monkeypatch.setattr( + inference_utils.torch.distributed, "is_initialized", lambda: distributed_initialized + ) + monkeypatch.setattr(inference_utils.torch.distributed, "get_world_size", lambda: world_size) + cuda_long_tensor = mock.Mock() + broadcast = mock.Mock() + monkeypatch.setattr(inference_utils.torch.cuda, "LongTensor", cuda_long_tensor) + monkeypatch.setattr(inference_utils.torch.distributed, "broadcast", broadcast) + + assert inference_utils.get_curr_time(do_broadcast=do_broadcast) == 123.0 + cuda_long_tensor.assert_not_called() + broadcast.assert_not_called() + + +def test_get_curr_time_broadcasts_for_multi_rank_sync(monkeypatch): + """Ensure explicit multi-rank timing still broadcasts a rank-zero timestamp.""" + timestamp = mock.Mock() + timestamp.item.return_value = 123_000_000_000 + monkeypatch.setattr(inference_utils.time, "time_ns", lambda: 123_000_000_000) + monkeypatch.setattr(inference_utils.torch.distributed, "is_initialized", lambda: True) + monkeypatch.setattr(inference_utils.torch.distributed, "get_world_size", lambda: 2) + monkeypatch.setattr(inference_utils.torch.cuda, "LongTensor", lambda _value: timestamp) + broadcast = mock.Mock() + monkeypatch.setattr(inference_utils.torch.distributed, "broadcast", broadcast) + + assert inference_utils.get_curr_time() == 123.0 + broadcast.assert_called_once_with(timestamp, src=0) diff --git a/tests/unit_tests/inference/test_data_parallel_inference_coordinator.py b/tests/unit_tests/inference/test_data_parallel_inference_coordinator.py index 8e9985b6dd5..e29d27a8cf0 100644 --- a/tests/unit_tests/inference/test_data_parallel_inference_coordinator.py +++ b/tests/unit_tests/inference/test_data_parallel_inference_coordinator.py @@ -2,6 +2,7 @@ import asyncio import itertools +import logging import multiprocessing import os import time @@ -18,6 +19,7 @@ from megatron.core.inference.data_parallel_inference_coordinator import ( DataParallelInferenceCoordinator, ) +from megatron.core.inference.data_parallel_inference_coordinator.handlers import handle_engine_reply from megatron.core.inference.engines.async_zmq_communicator import AsyncZMQCommunicator from megatron.core.inference.engines.dynamic_engine import ( DynamicInferenceEngine, @@ -776,22 +778,24 @@ def _make_routing_coordinator( class TestRoutingPolicies: """Unit tests for routing behavior under different policies and load conditions.""" - def test_no_prefix_caching_uses_round_robin(self): - """When prefix caching is off, round-robin is used regardless of load.""" + def test_no_prefix_caching_uses_load_balanced(self): + """When prefix caching is off, routing goes to the least-loaded rank.""" coord = _make_routing_coordinator(num_ranks=3, enable_prefix_caching=False) coord._pending_counts[coord.identity_to_rank_index[b"rank-0"]] = 2 coord._pending_counts[coord.identity_to_rank_index[b"rank-1"]] = 1 - results = [coord.get_best_data_parallel_rank([]) for _ in range(6)] - assert results == [b"rank-0", b"rank-1", b"rank-2", b"rank-0", b"rank-1", b"rank-2"] + # rank-2 has the fewest in-flight requests (0). + assert coord.get_best_data_parallel_rank([]) == b"rank-2" - def test_empty_hashes_uses_round_robin(self): - """Empty hash list falls back to round-robin.""" + def test_empty_hashes_uses_load_balanced(self): + """Empty hash list falls back to the least-loaded rank.""" coord = _make_routing_coordinator(num_ranks=4) + coord._pending_counts[coord.identity_to_rank_index[b"rank-0"]] = 3 coord._pending_counts[coord.identity_to_rank_index[b"rank-1"]] = 5 + coord._pending_counts[coord.identity_to_rank_index[b"rank-2"]] = 1 + coord._pending_counts[coord.identity_to_rank_index[b"rank-3"]] = 4 - results = [coord.get_best_data_parallel_rank([]) for _ in range(4)] - assert results == [b"rank-0", b"rank-1", b"rank-2", b"rank-3"] + assert coord.get_best_data_parallel_rank([]) == b"rank-2" def test_prefix_affinity_routing(self): """When prefix caching is on with hashes, scoring picks the best rank.""" @@ -846,18 +850,48 @@ def test_free_capacity_wins_when_prefix_rank_is_full(self): chosen = coord.get_best_data_parallel_rank([fake_hash]) assert chosen == b"rank-1" - def test_round_robin_policy_ignores_load(self): - """ROUND_ROBIN policy does naive round-robin regardless of load.""" + def test_load_balanced_policy_ignores_prefix(self): + """LOAD_BALANCED policy routes to the least-loaded rank, ignoring prefix affinity.""" coord = _make_routing_coordinator( num_ranks=3, enable_prefix_caching=True, - policy=PrefixCachingCoordinatorPolicy.ROUND_ROBIN, + policy=PrefixCachingCoordinatorPolicy.LOAD_BALANCED, ) - coord._pending_counts[coord.identity_to_rank_index[b"rank-0"]] = 1 + coord._pending_counts[coord.identity_to_rank_index[b"rank-0"]] = 2 coord._pending_counts[coord.identity_to_rank_index[b"rank-1"]] = 1 - coord._round_robin_idx = 0 - identities = list(coord.identities_of_data_parallel_ranks) - for i in range(len(identities)): - chosen = coord.get_best_data_parallel_rank([99]) - assert chosen == identities[i] + # Seed a prefix match on the most-loaded rank; load balancing must ignore it. + _set_hash_rank(coord, 99, b"rank-0", 1) + + assert coord.get_best_data_parallel_rank([99]) == b"rank-2" + + def test_reply_routing_survives_engine_removal(self, caplog): + """A removed engine's queued replies still deliver; never-connected senders assert.""" + + def reply(fid): + return [ + Headers.ENGINE_REPLY.value, + [{"request_id": fid, "generated_tokens": [1], "sampling_params": {}}], + ] + + coord = _make_routing_coordinator(num_ranks=2) + coord.tokenizer = DummyTokenizer() + coord.request_id_to_client_id = {11: b"client-A"} + coord.request_id_to_client_request_id = {11: 7} + coord.request_id_to_rank = {} + coord.router_socket = unittest.mock.MagicMock() + + # A sender that never registered is a protocol violation. + with pytest.raises(AssertionError, match="never-connected"): + handle_engine_reply(coord, b"impostor", reply(11)) + assert coord.router_socket.send_multipart.call_count == 0 + assert 11 in coord.request_id_to_client_id + + # Removal happens on failed *sends*, so the removed engine's in-flight + # reply can still arrive - and must reach its client. + coord._remove_engine(b"rank-0") + with caplog.at_level(logging.WARNING): + handle_engine_reply(coord, b"rank-0", reply(11)) + assert "removed engine" in caplog.text + assert coord.router_socket.send_multipart.call_args[0][0][0] == b"client-A" + assert 11 not in coord.request_id_to_client_id diff --git a/tests/unit_tests/inference/test_dynamic_prefix_caching_coordinator.py b/tests/unit_tests/inference/test_dynamic_prefix_caching_coordinator.py index b4b8da0e538..1d020ee1a79 100644 --- a/tests/unit_tests/inference/test_dynamic_prefix_caching_coordinator.py +++ b/tests/unit_tests/inference/test_dynamic_prefix_caching_coordinator.py @@ -351,23 +351,24 @@ def test_equal_scores_tiebreak_by_rank_index(self): selected = coordinator.get_best_data_parallel_rank(hashes) assert selected == rank_0 - def test_empty_hashes_uses_round_robin(self): - """Empty hash list falls back to round-robin.""" + def test_empty_hashes_uses_load_balanced(self): + """Empty hash list falls back to the least-loaded rank.""" coordinator = make_coordinator_direct() - for identity in coordinator.identities_of_data_parallel_ranks: - coordinator._pending_counts[coordinator.identity_to_rank_index[identity]] = 1 - rank1 = coordinator.get_best_data_parallel_rank([]) - rank2 = coordinator.get_best_data_parallel_rank([]) - assert rank1 != rank2 - - def test_disabled_prefix_caching_uses_round_robin(self): - """With prefix caching disabled, always uses round-robin.""" + identities = list(coordinator.identities_of_data_parallel_ranks) + for identity in identities: + coordinator._pending_counts[coordinator.identity_to_rank_index[identity]] = 2 + # Make the second rank the least loaded. + coordinator._pending_counts[coordinator.identity_to_rank_index[identities[1]]] = 0 + assert coordinator.get_best_data_parallel_rank([]) == identities[1] + + def test_disabled_prefix_caching_uses_load_balanced(self): + """With prefix caching disabled, always routes to the least-loaded rank.""" coordinator = make_coordinator_direct(enable_prefix_caching=False) - for identity in coordinator.identities_of_data_parallel_ranks: - coordinator._pending_counts[coordinator.identity_to_rank_index[identity]] = 1 - rank1 = coordinator.get_best_data_parallel_rank([1, 2, 3]) - rank2 = coordinator.get_best_data_parallel_rank([1, 2, 3]) - assert rank1 != rank2 + identities = list(coordinator.identities_of_data_parallel_ranks) + for identity in identities: + coordinator._pending_counts[coordinator.identity_to_rank_index[identity]] = 2 + coordinator._pending_counts[coordinator.identity_to_rank_index[identities[1]]] = 0 + assert coordinator.get_best_data_parallel_rank([1, 2, 3]) == identities[1] class TestCoordinatorShadowState: @@ -438,7 +439,7 @@ def test_routing_then_state_update_flow(self): tokens = [1, 2, 3, 4, 5, 6, 7, 8] hashes = coordinator.compute_request_hashes(tokens) - # First request: no matches, round-robin. + # First request: no matches, routed by load (least-loaded rank). rank = coordinator.get_best_data_parallel_rank(hashes) coordinator._update_rank_hashes(rank, hashes) diff --git a/tests/unit_tests/inference/test_hybrid_moe.py b/tests/unit_tests/inference/test_hybrid_moe.py index 96bb41903e9..a8cc9b743e6 100644 --- a/tests/unit_tests/inference/test_hybrid_moe.py +++ b/tests/unit_tests/inference/test_hybrid_moe.py @@ -164,6 +164,7 @@ def _build_context( use_cuda_graphs_for_non_decode_steps=True, max_requests=None, max_tokens=None, + cuda_graph_max_tokens=512, ): mamba_config = MambaInferenceStateConfig.from_model(model) return DynamicInferenceContext( @@ -178,6 +179,7 @@ def _build_context( use_cuda_graphs_for_non_decode_steps=use_cuda_graphs_for_non_decode_steps, max_requests=max_requests, max_tokens=max_tokens, + cuda_graph_max_tokens=cuda_graph_max_tokens, ), ) @@ -285,7 +287,7 @@ def test_nvls_ep_state_cross_product(self, rank_states): is_dummy = my_state == NONE model = self._build_model() - ctx = self._build_context(model, max_requests=64, max_tokens=512) + ctx = self._build_context(model, max_requests=64, max_tokens=512, cuda_graph_max_tokens=64) # Pre-capture every cuda graph in lockstep across EP ranks (mirrors # DynamicInferenceEngine.create_cuda_graphs in production). Without @@ -428,7 +430,10 @@ def test_nccl_eager_fallback_when_tokens_exceed_capacity(self, peer_state): # exceeded by a prefill-heavy rank. small_max_requests = 16 ctx = self._build_context( - model, use_cuda_graphs_for_non_decode_steps=True, max_requests=small_max_requests + model, + use_cuda_graphs_for_non_decode_steps=True, + max_requests=small_max_requests, + cuda_graph_max_tokens=small_max_requests, ) # Even EP ranks are dummy (no requests). Odd EP ranks get a state diff --git a/tests/unit_tests/inference/test_inference_config.py b/tests/unit_tests/inference/test_inference_config.py index d7e13ea3325..4f4dd780b64 100644 --- a/tests/unit_tests/inference/test_inference_config.py +++ b/tests/unit_tests/inference/test_inference_config.py @@ -26,8 +26,10 @@ def test_mutual_exclusivity_with_transformer_config(self): "async_sched_mode, expected", [ (None, AsyncScheduleMode.LEGACY), - ("serial", AsyncScheduleMode.SERIAL), - (AsyncScheduleMode.SERIAL, AsyncScheduleMode.SERIAL), + ("legacy", AsyncScheduleMode.LEGACY), + (AsyncScheduleMode.LEGACY, AsyncScheduleMode.LEGACY), + ("async", AsyncScheduleMode.ASYNC), + (AsyncScheduleMode.ASYNC, AsyncScheduleMode.ASYNC), ], ) def test_async_sched_mode_default_and_coercion(self, async_sched_mode, expected): @@ -35,16 +37,24 @@ def test_async_sched_mode_default_and_coercion(self, async_sched_mode, expected) kwargs = {} if async_sched_mode is None else {"async_sched_mode": async_sched_mode} assert InferenceConfig(**kwargs).async_sched_mode == expected - def test_async_sched_mode_rejects_invalid_value(self): + @pytest.mark.parametrize("invalid_mode", ["serial", "overlap", "invalid"]) + def test_async_sched_mode_rejects_invalid_value(self, invalid_mode): """Ensure invalid async scheduling modes fail during config construction.""" with pytest.raises(ValueError): - InferenceConfig(async_sched_mode="invalid") + InferenceConfig(async_sched_mode=invalid_mode) def test_async_sched_argparse_plumbing(self): """Ensure the CLI exposes async scheduling mode.""" parser = _add_inference_args(ArgumentParser()) - args = parser.parse_args(["--inference-dynamic-batching-async-sched-mode", "serial"]) - assert args.inference_dynamic_batching_async_sched_mode == "serial" + args = parser.parse_args(["--inference-dynamic-batching-async-sched-mode", "async"]) + assert args.inference_dynamic_batching_async_sched_mode == "async" + + @pytest.mark.parametrize("invalid_mode", ["serial", "overlap"]) + def test_async_sched_argparse_rejects_removed_modes(self, invalid_mode): + """Ensure the CLI rejects removed async scheduling modes.""" + parser = _add_inference_args(ArgumentParser()) + with pytest.raises(SystemExit): + parser.parse_args(["--inference-dynamic-batching-async-sched-mode", invalid_mode]) def test_inference_setup_config_maps_async_sched_mode(self): """Ensure declarative inference config maps async scheduling mode to runtime config.""" @@ -54,7 +64,36 @@ def test_inference_setup_config_maps_async_sched_mode(self): pg_collection="pg", decoder=SimpleNamespace(layer_type_list=None), ) - setup_config = InferenceSetupConfig(inference_dynamic_batching_async_sched_mode="serial") + setup_config = InferenceSetupConfig(inference_dynamic_batching_async_sched_mode="async") + + inference_config = setup_config.to_inference_config( + model=model, + kv_cache_management_mode="persist", + static_kv_memory_pointers=False, + enable_cuda_graphs=False, + verbose=False, + ) + + assert inference_config.async_sched_mode == AsyncScheduleMode.ASYNC + + def test_offset_sampling_seed_argparse_plumbing(self): + """Ensure the CLI can select a shared sampling seed across DP ranks.""" + parser = _add_inference_args(ArgumentParser()) + default_args = parser.parse_args([]) + assert default_args.offset_sampling_seed_by_dp_rank is True + + disabled_args = parser.parse_args(["--use-same-sampling-seed-across-dp-ranks"]) + assert disabled_args.offset_sampling_seed_by_dp_rank is False + + def test_inference_setup_config_maps_offset_sampling_seed_by_dp_rank(self): + """Ensure declarative inference config maps DP seed offset to runtime config.""" + model = SimpleNamespace( + position_embedding_type="rope", + max_sequence_length=4096, + pg_collection="pg", + decoder=SimpleNamespace(layer_type_list=None), + ) + setup_config = InferenceSetupConfig(offset_sampling_seed_by_dp_rank=False) inference_config = setup_config.to_inference_config( model=model, @@ -64,4 +103,4 @@ def test_inference_setup_config_maps_async_sched_mode(self): verbose=False, ) - assert inference_config.async_sched_mode == AsyncScheduleMode.SERIAL + assert inference_config.offset_sampling_seed_by_dp_rank is False diff --git a/tests/unit_tests/inference/test_inference_request.py b/tests/unit_tests/inference/test_inference_request.py index eafed6cecc6..559c20af18f 100644 --- a/tests/unit_tests/inference/test_inference_request.py +++ b/tests/unit_tests/inference/test_inference_request.py @@ -4,6 +4,7 @@ import msgpack import numpy as np +import pytest import torch from megatron.core.inference.inference_request import ( @@ -261,3 +262,59 @@ def test_dynamic_inference_request_serialize_strips_event_add_engine(): rec_out = DynamicInferenceRequestRecord.deserialize(rec.serialize()) assert rec_out.latency == 1.0 assert rec_out.requests[0].request_id == 7 + + +@pytest.mark.parametrize( + ("return_prompt_tokens", "expected_prompt_field"), + [ + (False, None), # default: prompt_tokens dropped from payload + (True, ("tensor", [1, 2, 3, 4])), # opt-in: prompt_tokens preserved + ], +) +def test_dynamic_inference_request_serialize_return_prompt_tokens( + return_prompt_tokens, expected_prompt_field +): + """DynamicInferenceRequest.serialize() reports prompt_length unconditionally + (the API uses it for `usage.prompt_tokens` on the response) and drops the + prompt_tokens tensor from the wire payload unless + SamplingParams.return_prompt_tokens is True. This is the load-bearing + wire-cost optimization for long agentic-RL prompts. The same call must + (a) leave self.prompt_tokens intact on the local instance — the drop is + wire-only — and (b) keep the routing_indices shape check honest, which + now relies on the saved prompt_len rather than self.prompt_tokens (which + is temporarily None during the drop).""" + sp = SamplingParams( + num_tokens_to_generate=5, termination_id=0, return_prompt_tokens=return_prompt_tokens + ) + prompt = torch.tensor([1, 2, 3, 4]) + # prompt_len=4 + generated=[10] → total_tokens=5 → routing_indices.shape[0] must be 4. + routing = np.zeros((4, 2, 1), dtype=np.int32) + req = _make_dynamic_request( + prompt_tokens=prompt, sampling_params=sp, generated_tokens=[10], routing_indices=routing + ) + + obj = req.serialize() + + # prompt_length is always populated (independent of the drop). + assert obj["prompt_length"] == 4 + # Payload either preserves the tensor wrapper or drops it (present but None). + assert obj["prompt_tokens"] == expected_prompt_field + # Local instance is unaffected — the drop is wire-only. + assert torch.equal(req.prompt_tokens, prompt) + # routing_indices survives the drop path (shape check would have crashed on + # the temporarily-None self.prompt_tokens if the fix used self.prompt_tokens). + assert isinstance(obj["routing_indices"], tuple) and obj["routing_indices"][0] == "ndarray" + + +def test_dynamic_inference_request_serialize_prompt_length_absent(): + """When prompt_tokens is None on the request, serialize() must not crash + (the drop path is guarded on `prompt_tokens is not None`) and prompt_length + must be reported as None. The DP coordinator can dispatch error/finish + records without prompt_tokens, so this path is real.""" + sp = SamplingParams(num_tokens_to_generate=1, termination_id=0) + req = DynamicInferenceRequest(request_id=99, prompt_tokens=None, sampling_params=sp) + + obj = req.serialize() + + assert obj["prompt_length"] is None + assert obj["prompt_tokens"] is None diff --git a/tests/unit_tests/inference/test_kv_transfer_backends.py b/tests/unit_tests/inference/test_kv_transfer_backends.py new file mode 100644 index 00000000000..0f32bc2e5c9 --- /dev/null +++ b/tests/unit_tests/inference/test_kv_transfer_backends.py @@ -0,0 +1,197 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import pytest +import torch + +from megatron.core.inference.disaggregation.ssm_reshard import SSMShardLayout, SSMStateDims +from megatron.core.inference.disaggregation.transfer_backends import base + + +def test_backend_registry_selects_by_explicit_name(): + assert base.construct_kv_transfer_backend_class("nixl").name == "nixl" + + try: + base.construct_kv_transfer_backend_class("unsupported") + except ValueError as exc: + assert "expected 'nixl'" in str(exc) + else: + raise AssertionError("unsupported backend should raise") + + +def test_ssm_geometry_uses_conv_and_recurrent_state_names(): + layout = SSMShardLayout( + global_rank=0, + tp_size=1, + tp_rank=0, + layer_start=0, + num_layers=1, + dims=SSMStateDims(nheads=2, headdim=4, d_state=5, ngroups=1, d_conv=3), + ) + memory_buffer = torch.zeros(1, 3, 2, 4, 5) + geometry = base.compute_buffer_geometry( + memory_buffer, + expected_num_blocks=3, + backend_name="test", + heads_per_partition=2, + ssm_layout=layout, + ssm_state_kind="recurrent", + ) + + metadata = base.export_geometry_meta(geometry, layout) + assert metadata["ssm_layout"]["dims"]["nheads"] == 2 + assert "mamba_layout" not in metadata + + with pytest.raises(ValueError, match="'conv' or 'recurrent'"): + base.compute_buffer_geometry( + memory_buffer, + expected_num_blocks=3, + backend_name="test", + heads_per_partition=2, + ssm_layout=layout, + ssm_state_kind="ssm", + ) + + +def test_nixl_direct_backend_exports_metadata_with_fake_agent(monkeypatch): + from megatron.core.inference.disaggregation.transfer_backends import nixl as nixl_mod + + class FakeAgent: + def __init__(self, name): + self.name = name + + def get_agent_metadata(self): + return b"agent-meta" + + def register_memory(self, tensor): + return ("reg", tuple(tensor.shape)) + + monkeypatch.setattr(nixl_mod, "_HAVE_NIXL", True) + monkeypatch.setattr(nixl_mod, "nixl_agent", FakeAgent) + + backend = nixl_mod.NixlTransferBackend( + "prefill", torch.zeros(2, 3, 5, dtype=torch.float32), expected_num_blocks=3 + ) + metadata = backend.export_meta() + + assert metadata["agent_name"] == "prefill" + assert metadata["bytes_per_slice"] == 20 + assert metadata["num_outer"] == 2 + assert metadata["num_blocks"] == 3 + assert metadata["blocks_axis"] == 1 + + +def test_nixl_begin_pull_blocks_uses_remote_metadata_with_fake_agent(monkeypatch): + from megatron.core.inference.disaggregation.transfer_backends import nixl as nixl_mod + + class FakeAgent: + def __init__(self, name): + self.name = name + self.transferred = False + + def get_agent_metadata(self): + return b"local" + + def register_memory(self, tensor): + return ("reg", tuple(tensor.shape)) + + def add_remote_agent(self, metadata): + assert metadata == b"remote" + return "peer" + + def get_xfer_descs(self, tuples, mem_type): + assert mem_type == "VRAM" + return tuples + + def initialize_xfer(self, op, local_desc, remote_desc, peer_id): + assert op == "READ" + assert peer_id == "peer" + return (local_desc, remote_desc) + + def transfer(self, xfer): + self.transferred = True + + def check_xfer_state(self, xfer): + assert self.transferred + return "DONE" + + monkeypatch.setattr(nixl_mod, "_HAVE_NIXL", True) + monkeypatch.setattr(nixl_mod, "nixl_agent", FakeAgent) + + backend = nixl_mod.NixlTransferBackend( + "decode", torch.zeros(2, 3, 5, dtype=torch.float32), expected_num_blocks=3 + ) + peer_meta = { + "agent_name": "prefill", + "agent_metadata_b64": "cmVtb3Rl", + "base_addr": 1234, + "bytes_per_slice": 20, + "num_outer": 2, + "outer_stride_bytes": 60, + "num_blocks": 3, + "device_id": 0, + "blocks_axis": 1, + } + backend.begin_pull_blocks(peer_meta, [1], [2]).wait() + + assert backend._agent.transferred is True + + +def test_nixl_begin_pull_blocks_returns_pollable_handle(monkeypatch): + from megatron.core.inference.disaggregation.transfer_backends import nixl as nixl_mod + + class FakeAgent: + def __init__(self, name): + self.name = name + self.transfers = 0 + self.polls = 0 + + def get_agent_metadata(self): + return b"local" + + def register_memory(self, tensor): + return ("reg", tuple(tensor.shape)) + + def add_remote_agent(self, metadata): + assert metadata == b"remote" + return "peer" + + def get_xfer_descs(self, tuples, mem_type): + assert mem_type == "VRAM" + return tuples + + def initialize_xfer(self, op, local_desc, remote_desc, peer_id): + assert op == "READ" + assert peer_id == "peer" + return {"local": local_desc, "remote": remote_desc} + + def transfer(self, xfer): + self.transfers += 1 + + def check_xfer_state(self, xfer): + self.polls += 1 + return "DONE" if self.polls >= 2 else "PENDING" + + monkeypatch.setattr(nixl_mod, "_HAVE_NIXL", True) + monkeypatch.setattr(nixl_mod, "nixl_agent", FakeAgent) + + backend = nixl_mod.NixlTransferBackend( + "decode", torch.zeros(2, 3, 5, dtype=torch.float32), expected_num_blocks=3 + ) + peer_meta = { + "agent_name": "prefill", + "agent_metadata_b64": "cmVtb3Rl", + "base_addr": 1234, + "bytes_per_slice": 20, + "num_outer": 2, + "outer_stride_bytes": 60, + "num_blocks": 3, + "device_id": 0, + "blocks_axis": 1, + } + + handle = backend.begin_pull_blocks(peer_meta, [1], [2]) + + assert backend._agent.transfers == 1 + assert backend._agent.polls == 0 + assert handle.poll() is False + assert handle.poll() is True diff --git a/tests/unit_tests/inference/test_moe_dispatching_and_routing.py b/tests/unit_tests/inference/test_moe_dispatching_and_routing.py index 49b5df613f7..5b21ab4c364 100644 --- a/tests/unit_tests/inference/test_moe_dispatching_and_routing.py +++ b/tests/unit_tests/inference/test_moe_dispatching_and_routing.py @@ -439,3 +439,85 @@ def test_cuda_graph_dispatch_combine(self, max_rank_tokens, seed): expected_combined = (global_hidden[start:end].float() * ep_size).bfloat16() torch.testing.assert_close(graph_combined, expected_combined, atol=0, rtol=0) + + +# ────────────────────────────────────────────────────────────────────── +# mask_routing_padding kernel +# ────────────────────────────────────────────────────────────────────── + +from megatron.core.transformer.moe.inference_routing_mask_kernel import ( # noqa: E402 + HAVE_TRITON, + mask_routing_padding, +) + +requires_triton_cuda = pytest.mark.skipif( + not HAVE_TRITON or not torch.cuda.is_available(), + reason="mask_routing_padding requires triton and CUDA", +) + + +@pytest.mark.internal +@requires_triton_cuda +class TestMaskRoutingPadding: + """Unit tests for the CUDA-graph padding-row routing mask. + + ``mask_routing_padding`` fills ``routing_map[real_token_count:, :]`` with -1 so + the NVLS dispatcher routes padding rows to no expert. ``real_token_count`` is in + the global (pre-SP-shard) frame; a non-zero ``tp_rank`` shifts local rows into + that frame before the comparison. Runs standalone — no context or NVLS hardware. + """ + + TOPK = 6 + + def _routing_map(self, n_rows, fill=3): + # All entries non-negative so masked (-1) slots are unambiguous. + return torch.full((n_rows, self.TOPK), fill, dtype=torch.int64, device="cuda") + + def _real_token_count(self, count): + return torch.tensor([count], dtype=torch.int32, device="cuda") + + @pytest.mark.parametrize("n_rows, real_count", [(16, 10), (128, 1), (7, 7), (64, 0)]) + def test_masks_rows_past_real_count(self, n_rows, real_count): + """Rows >= real_count become -1; rows < real_count are untouched (tp_rank=0).""" + routing_map = self._routing_map(n_rows) + original = routing_map.clone() + + mask_routing_padding(routing_map, self._real_token_count(real_count), tp_rank=0) + + torch.testing.assert_close(routing_map[:real_count], original[:real_count]) + assert torch.all(routing_map[real_count:] == -1) + + def test_real_count_equal_rows_is_noop(self): + """real_count == n_rows masks nothing (the unpadded decode case).""" + routing_map = self._routing_map(32) + original = routing_map.clone() + + mask_routing_padding(routing_map, self._real_token_count(32), tp_rank=0) + + torch.testing.assert_close(routing_map, original) + + def test_sp_rank_offset(self): + """Local rows are shifted by tp_rank * n_rows into the global frame. + + With 8 local rows on SP rank 1, local row r is global row r + 8. A global + real_count of 11 keeps global rows [8, 11) real (local [0, 3)) and masks + global rows [11, 16) (local [3, 8)). + """ + routing_map = self._routing_map(8) + original = routing_map.clone() + + mask_routing_padding(routing_map, self._real_token_count(11), tp_rank=1) + + torch.testing.assert_close(routing_map[:3], original[:3]) + assert torch.all(routing_map[3:] == -1) + + def test_sp_rank_fully_masked(self): + """An SP rank entirely beyond real_count is fully masked. + + 8 local rows on rank 1 cover global rows [8, 16); real_count=8 masks all. + """ + routing_map = self._routing_map(8) + + mask_routing_padding(routing_map, self._real_token_count(8), tp_rank=1) + + assert torch.all(routing_map == -1) diff --git a/tests/unit_tests/inference/test_mtp_cuda_graph_inference.py b/tests/unit_tests/inference/test_mtp_cuda_graph_inference.py index 8f738ceb81c..0a1f70dcde3 100644 --- a/tests/unit_tests/inference/test_mtp_cuda_graph_inference.py +++ b/tests/unit_tests/inference/test_mtp_cuda_graph_inference.py @@ -1096,3 +1096,366 @@ def test_nccl_ep_dummy_bailout_with_decode_only_cuda_graphs(self, peer_state): assert h_out.shape == (tp_size, 1, self.HIDDEN_SIZE) assert logits.shape == (tp_size, 1, self.VOCAB_SIZE) + + +# --------------------------------------------------------------------------- # +# TestMTPBlockScopeCudaGraph (TP = 1) +# --------------------------------------------------------------------------- # + + +def _build_hybrid_stack_spec(): + """Build a minimal HybridStack spec using local (non-TE) modules.""" + attention_layer_spec = get_gpt_layer_local_spec() + + backend = LocalSpecProvider() + norm_impl = backend.layer_norm() + col_linear_impl = backend.column_parallel_linear() + + mtp_layer_spec = ModuleSpec( + module=MultiTokenPredictionLayer, + submodules=MultiTokenPredictionLayerSubmodules( + enorm=norm_impl, + hnorm=norm_impl, + eh_proj=col_linear_impl, + mtp_model_layer=None, + layer_norm=norm_impl, + ), + ) + mtp_block_spec = ModuleSpec( + module=MultiTokenPredictionBlock, + submodules=MultiTokenPredictionBlockSubmodules(layer_specs=[mtp_layer_spec]), + ) + + return ModuleSpec( + module=HybridStack, + submodules=HybridStackSubmodules( + attention_layer=attention_layer_spec, mtp_block_spec=mtp_block_spec + ), + ) + + +class TestMTPBlockScopeCudaGraph: + """Tests that block-scope CUDA graphs correctly propagate decoder hidden + states for MTP inference. + + When ``inference_cuda_graph_scope='block'``, the entire model forward is + captured as a single CUDA graph. ``context.mtp_decoder_hidden_states`` + holds a pre-allocated buffer that is written via ``copy_()`` during every + graph replay so that it is available on every replay, not just during capture. + """ + + HIDDEN_SIZE = 32 + VOCAB_SIZE = 100 + MAX_SEQ_LEN = 64 + NUM_LAYERS = 4 + NUM_ATTN_HEADS = 4 + + @classmethod + def setup_class(cls): + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, pipeline_model_parallel_size=1 + ) + + @classmethod + def teardown_class(cls): + delete_cuda_graphs() + Utils.destroy_model_parallel() + + def teardown_method(self): + delete_cuda_graphs() + + def _build_model(self, *, inference_cuda_graph_scope='block', model_type='hybrid'): + """Build a GPT or Hybrid model with MTP and local CUDA graph support.""" + model_parallel_cuda_manual_seed(123, inference_rng_tracker=True, force_reset_rng=True) + config = TransformerConfig( + num_layers=self.NUM_LAYERS, + hidden_size=self.HIDDEN_SIZE, + num_attention_heads=self.NUM_ATTN_HEADS, + use_cpu_initialization=True, + attention_backend=AttnBackend.local, + params_dtype=torch.bfloat16, + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + pipeline_dtype=torch.bfloat16, + mtp_num_layers=1, + cuda_graph_impl="local", + inference_cuda_graph_scope=inference_cuda_graph_scope, + ) + if model_type == 'gpt': + layer_spec = get_gpt_layer_local_spec() + mtp_block_spec = get_gpt_mtp_block_spec( + config=config, spec=layer_spec, use_transformer_engine=False + ) + model = GPTModel( + config=config, + transformer_layer_spec=layer_spec, + mtp_block_spec=mtp_block_spec, + vocab_size=self.VOCAB_SIZE, + max_sequence_length=self.MAX_SEQ_LEN, + parallel_output=True, + pre_process=True, + post_process=True, + position_embedding_type='rope', + ).cuda() + elif model_type == 'hybrid': + hybrid_stack_spec = _build_hybrid_stack_spec() + model = HybridModel( + config=config, + hybrid_stack_spec=hybrid_stack_spec, + vocab_size=self.VOCAB_SIZE, + max_sequence_length=self.MAX_SEQ_LEN, + parallel_output=True, + pre_process=True, + post_process=True, + hybrid_layer_pattern="****/*", + position_embedding_type='rope', + ).cuda() + else: + raise ValueError(f"Unknown model_type: {model_type!r}") + for param in model.parameters(): + param.data = param.data.to(config.params_dtype) + model.eval() + return model + + def _build_engine( + self, *, inference_cuda_graph_scope='block', num_speculative_tokens=1, model_type='hybrid' + ): + """Build a DynamicInferenceEngine with block-scope CUDA graphs.""" + delete_cuda_graphs() + model = self._build_model( + inference_cuda_graph_scope=inference_cuda_graph_scope, model_type=model_type + ) + config = model.config + context = DynamicInferenceContext( + model_config=config, + inference_config=InferenceConfig( + max_sequence_length=self.MAX_SEQ_LEN, + buffer_size_gb=0.5, + materialize_only_last_token_logits=False, + num_speculative_tokens=num_speculative_tokens, + block_size_tokens=256, + max_requests=16, + num_cuda_graphs=-1, + sampling_backend='torch', + ), + ) + wrapped = GPTInferenceWrapper(model, context) + wrapped.model_is_pipeline_parallel = False + mock_tokenizer = mock.Mock() + ctrl = TextGenerationController(inference_wrapped_model=wrapped, tokenizer=mock_tokenizer) + engine = DynamicInferenceEngine(ctrl, context) + return engine + + @pytest.mark.parametrize("model_type", ['gpt', 'hybrid']) + @pytest.mark.parametrize("inference_cuda_graph_scope", ['block', 'layer']) + @torch.inference_mode() + def test_decoder_hidden_states_set_after_forward(self, inference_cuda_graph_scope, model_type): + """Decoder hidden states are accessible via the context after each forward pass. + + Block-scope CUDA graphs: forward() writes via copy_() into the pre-allocated + context buffer, captured once and replayed to the same GPU address each step. + Layer-scope (non-block) CUDA graphs: forward() assigns the tensor directly to + the context attribute; the controller sets it back to None after reading to allow GC. + Both scopes are valid with cuda_graph_impl='local'. Covers GPTModel and HybridModel. + """ + engine = self._build_engine( + inference_cuda_graph_scope=inference_cuda_graph_scope, model_type=model_type + ) + ctrl = engine.controller + context = engine.context + + prompt_length = 10 + req = DynamicInferenceRequest( + request_id=0, + prompt_tokens=torch.arange(prompt_length, device='cuda'), + sampling_params=SamplingParams(num_tokens_to_generate=20), + ) + context.add_request(req) + context.initialize_attention_state() + + active_mask = torch.ones(1, device='cuda', dtype=torch.int32) + new_tokens = torch.zeros(1, device='cuda', dtype=torch.int64) + new_spec = torch.zeros(1, 1, device='cuda', dtype=torch.int64) + context.update_requests( + active_requests_mask=active_mask, new_tokens=new_tokens, new_speculative_tokens=new_spec + ) + context.initialize_attention_state() + + for step in range(3): + # Simulate the controller consuming hidden states from the previous step. + if inference_cuda_graph_scope != 'block': + context.mtp_decoder_hidden_states = None + + input_ids, position_ids, _ = ctrl._dynamic_step_context_init() + ctrl._dynamic_step_forward_logits(input_ids, position_ids) + + assert context.mtp_decoder_hidden_states is not None, ( + f"Step {step}: context.mtp_decoder_hidden_states is None " + f"(scope={inference_cuda_graph_scope})" + ) + assert context.mtp_decoder_hidden_states.shape[-1] == self.HIDDEN_SIZE + + @torch.inference_mode() + def test_mtp_forward_with_runtime_tokens_below_max(self): + """Block-scope MTP forward is correct when runtime tokens < ``max_tokens``. + + The decoder hidden-states buffer is pre-allocated to the worst-case + ``(max_tokens, 1, hidden_size)`` so the block CUDA graph can write into a + fixed GPU address, but a real decode step only fills the ``[:n]`` prefix + (``n = active_request_count`` rows, far below ``max_tokens``). + + This verifies two things in one shot: + 1. **No shape issues** — the controller gathers ``[:n]`` rows from the + ``max_tokens``-sized buffer and the downstream MTP layer forward runs + without any shape mismatch. + 2. **Forward correctness** — the unused buffer tail does not leak into the + computation. The tail is poisoned with NaN; the sampled MTP tokens must + still match a reference run against an exactly ``(n, 1, H)``-sized + buffer (and stay finite/in-range). If the MTP forward read past the + valid prefix, the NaNs would change the result. + """ + engine = self._build_engine(inference_cuda_graph_scope='block') + ctrl = engine.controller + context = engine.context + + num_spec = ctrl.num_speculative_tokens + assert num_spec > 0 and ctrl.num_mtp_depths > 0 + + # The block-scope buffer is sized to the worst case; pick an active count + # that is comfortably below capacity to exercise the partial-fill path. + active_request_count = 3 + assert active_request_count < context.max_tokens + + buffer = context.mtp_decoder_hidden_states + assert buffer is not None + assert buffer.shape == (context.max_tokens, 1, self.HIDDEN_SIZE) + + # Deterministic hidden states for the valid prefix, shared by both runs. + torch.manual_seed(42) + prefix_hidden = torch.randn( + active_request_count, 1, self.HIDDEN_SIZE, device='cuda', dtype=torch.bfloat16 + ) + + def _run_eager_mtp(decoder_hidden_states): + """Set up decode state and run eager MTP, returning sampled tokens.""" + context.reset() + context.total_request_count = active_request_count + context.paused_request_count = 0 + context.request_kv_length_offsets[:active_request_count] = torch.arange( + active_request_count, dtype=torch.int32, device='cuda' + ) + context.request_query_lengths[:active_request_count] = torch.ones( + active_request_count, dtype=torch.int32, device='cuda' + ) + + ctrl.num_speculative_tokens = num_spec + ctrl._init_mtp_sampling_tensors() + ctrl._mtp_token_ids_buf.zero_() + ctrl._mtp_position_ids_buf.zero_() + ctrl._sampled_tokens_cuda[:active_request_count] = torch.remainder( + torch.arange(active_request_count, device='cuda'), self.VOCAB_SIZE + ) + + # Eager path (no CUDA graph, no SP padding for TP=1). + ctrl._mtp_resolved_padded_count = None + context._using_cuda_graph_this_step = False + + context.mtp_decoder_hidden_states = decoder_hidden_states + ctrl._last_accepted_seq_indices = torch.arange(active_request_count, device='cuda') + + # Greedy sampling for all active requests. + context.active_request_metadata["temperature"][:active_request_count] = 1.0 + context.active_request_metadata["top_k"][:active_request_count] = 1 + context.active_request_metadata["top_p"][:active_request_count] = 0.0 + + ctrl._compute_serial_mtp_and_sample() + + return [ + ctrl._sampled_mtp_tokens_cuda[d, :active_request_count].clone() + for d in range(ctrl.num_mtp_depths) + ] + + # Run 1: max_tokens-sized buffer with only the [:n] prefix valid; poison + # the unused tail with NaN so any over-read corrupts the result. + buffer.fill_(float('nan')) + buffer[:active_request_count].copy_(prefix_hidden) + oversized_tokens = _run_eager_mtp(buffer) + + # Run 2: reference buffer sized exactly to the runtime token count. + exact_buffer = prefix_hidden.clone() + reference_tokens = _run_eager_mtp(exact_buffer) + + for depth in range(ctrl.num_mtp_depths): + sampled = oversized_tokens[depth] + assert sampled.shape == (active_request_count,), ( + f"depth={depth}: expected shape ({active_request_count},), " + f"got {tuple(sampled.shape)}" + ) + assert sampled.dtype == torch.int64 + assert torch.all(sampled >= 0) and torch.all(sampled < self.VOCAB_SIZE) + assert torch.equal(sampled, reference_tokens[depth]), ( + f"depth={depth}: MTP tokens from the max_tokens-sized buffer " + f"{sampled.tolist()} != reference {reference_tokens[depth].tolist()}; " + "the unused buffer tail leaked into the MTP forward" + ) + + @pytest.mark.parametrize("model_type", ['gpt', 'hybrid']) + @pytest.mark.parametrize("inference_cuda_graph_scope", ['block', 'layer']) + @torch.inference_mode() + def test_no_spec_decode_leaves_decoder_hidden_states_unset( + self, inference_cuda_graph_scope, model_type + ): + """Regression: a model with an MTP head but ``num_speculative_tokens == 0``. + + When the model has MTP layers (``mtp_num_layers >= 1``) but speculative + decoding is disabled, plain inference must NOT touch + ``context.mtp_decoder_hidden_states`` — there is no serial post-verification + MTP step to consume it, and for block-scope CUDA graphs the buffer is never + even allocated (it is allocated only when ``num_speculative_tokens > 0``). + + Covers both GPTModel and HybridModel since each carries the same MTP + post-process branch (``gpt_model.py`` / ``hybrid_model.py``). + """ + engine = self._build_engine( + inference_cuda_graph_scope=inference_cuda_graph_scope, + num_speculative_tokens=0, + model_type=model_type, + ) + ctrl = engine.controller + context = engine.context + + # No speculative decoding -> no MTP depths and no pre-allocated buffer. + assert ctrl.num_speculative_tokens == 0 + assert ctrl.num_mtp_depths == 0 + assert context.mtp_decoder_hidden_states is None + + prompt_length = 10 + req = DynamicInferenceRequest( + request_id=0, + prompt_tokens=torch.arange(prompt_length, device='cuda'), + sampling_params=SamplingParams(num_tokens_to_generate=20), + ) + context.add_request(req) + context.initialize_attention_state() + + active_mask = torch.ones(1, device='cuda', dtype=torch.int32) + new_tokens = torch.zeros(1, device='cuda', dtype=torch.int64) + context.update_requests( + active_requests_mask=active_mask, new_tokens=new_tokens, new_speculative_tokens=None + ) + context.initialize_attention_state() + + # Force the inference flag on so the forward takes the in_inference_mode + # branch even though we drive the step directly rather than via the engine + # run loop. + with InferenceMode.active(): + for step in range(3): + input_ids, position_ids, _ = ctrl._dynamic_step_context_init() + ctrl._dynamic_step_forward_logits(input_ids, position_ids) + + assert context.mtp_decoder_hidden_states is None, ( + f"Step {step}: mtp_decoder_hidden_states should stay None when " + f"num_speculative_tokens == 0 (model={model_type}, " + f"scope={inference_cuda_graph_scope}), got a tensor of shape " + f"{tuple(context.mtp_decoder_hidden_states.shape)}" + ) diff --git a/tests/unit_tests/inference/test_nccl_transfer_backend.py b/tests/unit_tests/inference/test_nccl_transfer_backend.py new file mode 100644 index 00000000000..ecb5fb44b27 --- /dev/null +++ b/tests/unit_tests/inference/test_nccl_transfer_backend.py @@ -0,0 +1,138 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Distributed unit test of the two-sided NCCL transfer backend. + +Prefill TP2 {0,1} -> decode TP1 {2} on real GPUs: the decode posts +begin_pull_blocks, the prefills post the matching begin_push_blocks, and the +decode's paged buffer must end up byte-identical to a direct shard of a known +global KV. Exercises the hetero head-merge through the same reshard plan the +NIXL backend uses. The test uses the process group provided by the unit-test +runner instead of spawning a nested distributed job. +""" + +import os + +import pytest +import torch + +L, H, HD, T, NB = 4, 8, 16, 8, 6 # layers, kv heads, head dim, tokens/block, pool blocks +BLOCKS = [1, 3] # the request's blocks (same ids both sides for simplicity) + + +def _global_blocks(): + """Global KV for the request's blocks: (block, kv, layer, token, head, dim) + with a distinct value per (block, kv, layer, head).""" + g = torch.zeros(len(BLOCKS), 2, L, T, H, HD) + for b in range(len(BLOCKS)): + for kv in range(2): + for l in range(L): + for h in range(H): + g[b, kv, l, :, h, :] = ((b * 2 + kv) * L + l) * 100 + h + return g + + +def _backend(rank, tp_size, tp_rank, device): + from megatron.core.inference.disaggregation.transfer_backends.nccl import NcclTransferBackend + + heads_local = H // tp_size + buf = torch.zeros(2, L, NB, T, heads_local, HD, device=device) + backend = NcclTransferBackend( + agent_name=f"test-rank{rank}", + memory_buffer=buf, + expected_num_blocks=NB, + tp_size=tp_size, + tp_rank=tp_rank, + num_kv_heads_global=H, + heads_per_partition=heads_local, + head_dim=HD, + tokens_per_block=T, + global_rank=rank, + pp_size=1, + pp_rank=0, + num_layers_global=L, + layer_start=0, + layer_end=L, + ) + return backend, buf + + +def _meta_stub(rank, tp_size, tp_rank): + """A rank's export_meta, built without its backend; the address fields are + unused by NCCL and the geometry is deterministic.""" + heads_local = H // tp_size + return { + "transport": "nccl", + "nccl_rank": rank, + "num_blocks": NB, + "blocks_axis": 2, + "num_outer": 2 * L, + "heads_per_partition": heads_local, + "head_dim": HD, + "tokens_per_block": T, + "element_size": 4, + "bytes_per_slice": T * heads_local * HD * 4, + "outer_stride_bytes": NB * T * heads_local * HD * 4, + "base_addr": 0, + "device_id": rank, + "global_rank": rank, + "tp_size": tp_size, + "tp_rank": tp_rank, + "pp_size": 1, + "pp_rank": 0, + "num_layers_global": L, + "num_kv_heads_global": H, + "layer_start": 0, + "layer_end": L, + } + + +@pytest.mark.skipif( + not ( + torch.cuda.is_available() + and torch.cuda.device_count() >= 3 + and int(os.environ.get("WORLD_SIZE", "1")) >= 3 + ), + reason="requires torchrun with >=3 CUDA ranks (prefill TP2 {0,1} + decode TP1 {2})", +) +def test_nccl_push_pull_tp2_to_tp1(): + import torch.distributed as dist + + local_rank = int(os.environ["LOCAL_RANK"]) + torch.cuda.set_device(local_rank) + device = f"cuda:{local_rank}" + if not dist.is_initialized(): + dist.init_process_group("nccl") + + rank = dist.get_rank() + control_group = dist.new_group(backend="gloo") + + # Initialize the default NCCL communicator collectively before ranks 0–2 + # use it for point-to-point transfers. Extra CI ranks synchronize through + # the Gloo control group and do not issue conflicting NCCL collectives. + dist.barrier() + + transfer_ok = True + g = _global_blocks().to(device) + if rank in (0, 1): # prefill TP2 + backend, buf = _backend(rank, 2, rank, device) + heads = slice(rank * (H // 2), (rank + 1) * (H // 2)) + for i, block in enumerate(BLOCKS): + # buffer layout [2, L, B, T, h, d] + buf[:, :, block] = g[i, :, :, :, heads, :] + # In production the decode's metas arrive in SEND_KV. + handle = backend.begin_push_blocks({"tp_metas": [_meta_stub(2, 1, 0)]}, BLOCKS) + handle.wait() + elif rank == 2: # decode TP1 + backend, buf = _backend(rank, 1, 0, device) + # In production the prefills' metas arrive in the hand-off kv_meta. + metas = [_meta_stub(0, 2, 0), _meta_stub(1, 2, 1)] + handle = backend.begin_pull_blocks({"tp_metas": metas}, BLOCKS, BLOCKS) + handle.wait() + expected = torch.zeros_like(buf) + for i, block in enumerate(BLOCKS): + expected[:, :, block] = g[i] + transfer_ok = torch.equal(buf, expected) + + result = torch.tensor(int(transfer_ok)) + dist.all_reduce(result, op=dist.ReduceOp.MIN, group=control_group) + assert result.item() == 1 diff --git a/tests/unit_tests/inference/test_mamba_reshard.py b/tests/unit_tests/inference/test_ssm_reshard.py similarity index 66% rename from tests/unit_tests/inference/test_mamba_reshard.py rename to tests/unit_tests/inference/test_ssm_reshard.py index 4a197813ab9..584f3a0213e 100644 --- a/tests/unit_tests/inference/test_mamba_reshard.py +++ b/tests/unit_tests/inference/test_ssm_reshard.py @@ -1,10 +1,10 @@ # Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. -"""Hetero TP/PP reshard of Mamba conv/ssm state (pure, CPU). +"""Hetero TP/PP reshard of SSM conv/recurrent state (pure, CPU). -Builds a known global Mamba state, shards it to a source (tp,pp) the exact way -mamba_mixer does ([x|B|C] conv bands + head-sharded ssm, layers split by PP), -runs plan_mamba_reshard to a different destination (tp,pp), and asserts every +Builds a known global SSM state, shards it to a source (tp,pp) the exact way +MambaMixer does ([x|B|C] conv bands + head-sharded recurrent state, layers +split by PP), runs ``plan_ssm_reshard`` to a different destination (tp,pp), and asserts every destination rank ends up byte-identical to a direct shard of the global state. This validates the band/layer index math against the real sharding model without a hybrid checkpoint (the residual gap is a real-model functional run). @@ -13,10 +13,10 @@ import pytest import torch -from megatron.core.inference.disaggregation.mamba_reshard import ( - MambaShardLayout, - MambaStateDims, - plan_mamba_reshard, +from megatron.core.inference.disaggregation.ssm_reshard import ( + SSMShardLayout, + SSMStateDims, + plan_ssm_reshard, ) @@ -26,17 +26,17 @@ def apply_conv_transfer(t, src_conv, dst_conv): dst_conv[t.dst_layer, t.dst_lo : t.dst_hi, :] = src_conv[t.src_layer, t.src_lo : t.src_hi, :] -def apply_ssm_transfer(t, src_ssm, dst_ssm): - """Copy an ssm sub-block in-memory; ssm is +def apply_recurrent_transfer(t, src_recurrent, dst_recurrent): + """Copy a recurrent-state sub-block in-memory; recurrent state is ``(num_layers, nheads_local, headdim, d_state)`` -- the band slices heads.""" - dst_ssm[t.dst_layer, t.dst_lo : t.dst_hi, :, :] = src_ssm[ + dst_recurrent[t.dst_layer, t.dst_lo : t.dst_hi, :, :] = src_recurrent[ t.src_layer, t.src_lo : t.src_hi, :, : ] # Global model dims (chosen divisible by the tp values under test). NHEADS, HEADDIM, DSTATE, NGROUPS, DCONV = 8, 4, 2, 2, 3 -M = 4 # global Mamba layers +M = 4 # global SSM layers D_INNER = NHEADS * HEADDIM # 32 G = NGROUPS * DSTATE # 4 (B and C band global size) CONV_DIM = D_INNER + 2 * G # 40 @@ -45,38 +45,38 @@ def apply_ssm_transfer(t, src_ssm, dst_ssm): def _global_state(): """Distinct value per (layer, channel, ...) so any mis-slice is caught.""" conv = torch.arange(M * CONV_DIM * DCONV, dtype=torch.float32).reshape(M, CONV_DIM, DCONV) - ssm = ( + recurrent = ( torch.arange(M * NHEADS * HEADDIM * DSTATE, dtype=torch.float32).reshape( M, NHEADS, HEADDIM, DSTATE ) + 10_000.0 ) - return conv, ssm + return conv, recurrent def _layouts(tp, pp): - """One MambaShardLayout per rank for a (tp, pp) instance; rank = p*tp + r. + """One SSMShardLayout per rank for a (tp, pp) instance; rank = p*tp + r. PP splits the M layers evenly (contiguous per stage).""" per = M // pp out = {} for p in range(pp): for r in range(tp): rank = p * tp + r - out[rank] = MambaShardLayout( + out[rank] = SSMShardLayout( global_rank=rank, tp_size=tp, tp_rank=r, layer_start=p * per, num_layers=per, - dims=MambaStateDims( + dims=SSMStateDims( nheads=NHEADS, headdim=HEADDIM, d_state=DSTATE, ngroups=NGROUPS, d_conv=DCONV ), ) return out -def _shard(conv_g, ssm_g, lay: MambaShardLayout): - """Shard the global state to one rank exactly as mamba_mixer does.""" +def _shard(conv_g, recurrent_g, lay: SSMShardLayout): + """Shard the global state to one rank exactly as MambaMixer does.""" s, e = lay.layer_range() r, tp = lay.tp_rank, lay.tp_size di_l = D_INNER // tp @@ -86,8 +86,8 @@ def _shard(conv_g, ssm_g, lay: MambaShardLayout): c = conv_g[s:e, D_INNER + G : D_INNER + 2 * G][:, r * g_l : (r + 1) * g_l] conv_l = torch.cat([x, b, c], dim=1).contiguous() nh_l = NHEADS // tp - ssm_l = ssm_g[s:e, r * nh_l : (r + 1) * nh_l, :, :].contiguous() - return conv_l, ssm_l + recurrent_l = recurrent_g[s:e, r * nh_l : (r + 1) * nh_l, :, :].contiguous() + return conv_l, recurrent_l @pytest.mark.parametrize( @@ -101,12 +101,12 @@ def _shard(conv_g, ssm_g, lay: MambaShardLayout): ((2, 1), (2, 1)), # identity ], ) -def test_mamba_reshard_reconstructs_destination(src, dst): - conv_g, ssm_g = _global_state() +def test_ssm_reshard_reconstructs_destination(src, dst): + conv_g, recurrent_g = _global_state() src_lay, dst_lay = _layouts(*src), _layouts(*dst) # Source per-rank tensors (as a prefill instance would hold them). - src_t = {rk: _shard(conv_g, ssm_g, lay) for rk, lay in src_lay.items()} + src_t = {rk: _shard(conv_g, recurrent_g, lay) for rk, lay in src_lay.items()} # Destination buffers, zero-filled at each rank's local shape. dst_t = {} for rk, lay in dst_lay.items(): @@ -115,71 +115,73 @@ def test_mamba_reshard_reconstructs_destination(src, dst): torch.zeros(lay.num_layers, lay.nheads_local, HEADDIM, DSTATE), ) - plan = plan_mamba_reshard(list(src_lay.values()), list(dst_lay.values())) + plan = plan_ssm_reshard(list(src_lay.values()), list(dst_lay.values())) for t in plan: if t.is_conv: apply_conv_transfer(t, src_t[t.src_rank][0], dst_t[t.dst_rank][0]) else: - apply_ssm_transfer(t, src_t[t.src_rank][1], dst_t[t.dst_rank][1]) + apply_recurrent_transfer(t, src_t[t.src_rank][1], dst_t[t.dst_rank][1]) # Every destination rank must match a direct shard of the global state. for rk, lay in dst_lay.items(): - want_conv, want_ssm = _shard(conv_g, ssm_g, lay) + want_conv, want_recurrent = _shard(conv_g, recurrent_g, lay) assert torch.equal(dst_t[rk][0], want_conv), f"conv mismatch at rank {rk} ({src}->{dst})" - assert torch.equal(dst_t[rk][1], want_ssm), f"ssm mismatch at rank {rk} ({src}->{dst})" + assert torch.equal( + dst_t[rk][1], want_recurrent + ), f"recurrent mismatch at rank {rk} ({src}->{dst})" -def test_mamba_rejects_indivisible_groups(): +def test_ssm_rejects_indivisible_groups(): """ngroups < tp_size would truncate the B/C bands to zero width; reject it up front instead of silently dropping state.""" with pytest.raises(ValueError): - MambaShardLayout( + SSMShardLayout( global_rank=0, tp_size=4, tp_rank=0, layer_start=0, num_layers=1, - dims=MambaStateDims(nheads=8, headdim=HEADDIM, d_state=DSTATE, ngroups=2, d_conv=DCONV), + dims=SSMStateDims(nheads=8, headdim=HEADDIM, d_state=DSTATE, ngroups=2, d_conv=DCONV), ) -def test_mamba_dedupes_replica_sources(): - """Two source ranks holding the same Mamba shard (same tp_rank+layer_start, +def test_ssm_dedupes_replica_sources(): + """Two source ranks holding the same SSM shard (same tp_rank+layer_start, e.g. EP/DP replicas) are deduped: the shard is sourced from exactly one of them (smallest global_rank), so no duplicate sends.""" def _lay(gr): - return MambaShardLayout( + return SSMShardLayout( global_rank=gr, tp_size=1, tp_rank=0, layer_start=0, num_layers=M, - dims=MambaStateDims( + dims=SSMStateDims( nheads=NHEADS, headdim=HEADDIM, d_state=DSTATE, ngroups=NGROUPS, d_conv=DCONV ), ) - plan = plan_mamba_reshard([_lay(0), _lay(1)], [_lay(2)]) + plan = plan_ssm_reshard([_lay(0), _lay(1)], [_lay(2)]) assert {t.src_rank for t in plan} == {0} # only the smallest-rank replica sources def test_layout_wire_roundtrip(): - """Layouts cross the coordinator as plain dicts (asdict) and are rebuilt via - MambaShardLayout(**dict); the nested dims dict must coerce back to - MambaStateDims so proxies (.headdim/.d_conv/...) keep working.""" + """Layouts cross the coordinator as plain dicts (asdict) and are rebuilt + via SSMShardLayout(**dict); the nested dims dict must coerce back to + SSMStateDims.""" import dataclasses - lay = MambaShardLayout( + lay = SSMShardLayout( global_rank=1, tp_size=2, tp_rank=1, layer_start=0, num_layers=M, - dims=MambaStateDims( + dims=SSMStateDims( nheads=NHEADS, headdim=HEADDIM, d_state=DSTATE, ngroups=NGROUPS, d_conv=DCONV ), ) - rebuilt = MambaShardLayout(**dataclasses.asdict(lay)) + rebuilt = SSMShardLayout(**dataclasses.asdict(lay)) assert rebuilt == lay - assert rebuilt.headdim == HEADDIM and rebuilt.d_conv == DCONV + assert rebuilt.conv_dim_local == lay.conv_dim_local diff --git a/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py b/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py index 1bc3cda149f..e8a6254a6ef 100644 --- a/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py +++ b/tests/unit_tests/inference/text_generation_controllers/test_text_generation_controller.py @@ -33,7 +33,9 @@ ) from megatron.core.inference.sampling_params import SamplingParams from megatron.core.inference.text_generation_controllers.text_generation_controller import ( - DecodeForwardPrimer, + AsyncScheduleLogitsState, + DecodeOnly, + DynamicBatchControllerStepResult, TextGenerationController, ) from megatron.core.inference.utils import InferenceMode @@ -196,9 +198,18 @@ def setup_model( def _make_async_sched_context(total_request_count=2, paused_request_count=0): metadata_len = max(total_request_count, 1) - return SimpleNamespace( + request_metadata = { + "temperature": torch.ones(metadata_len), + "top_k": torch.ones(metadata_len, dtype=torch.int64), + "top_p": torch.zeros(metadata_len), + "return_log_probs": torch.zeros(metadata_len, dtype=torch.bool), + "top_n_logprobs": torch.zeros(metadata_len, dtype=torch.int64), + "skip_prompt_log_probs": torch.zeros(metadata_len, dtype=torch.bool), + "termination_id": torch.full((metadata_len,), 99, dtype=torch.int64), + } + context = SimpleNamespace( config=SimpleNamespace( - materialize_only_last_token_logits=True, async_sched_mode=AsyncScheduleMode.SERIAL + materialize_only_last_token_logits=True, async_sched_mode=AsyncScheduleMode.ASYNC ), is_hybrid_model=False, enable_prefix_caching=False, @@ -208,14 +219,22 @@ def _make_async_sched_context(total_request_count=2, paused_request_count=0): chunked_prefill_request_id=-1, num_prefill_requests=0, padded_active_request_count=8, + step_count=0, + prefix_cache_lru_clock=0, + lifetime_prefill_token_count=0, request_ids=torch.arange(10, 10 + metadata_len, dtype=torch.int32), - request_metadata={ - "top_k": torch.ones(metadata_len, dtype=torch.int64), - "top_p": torch.zeros(metadata_len), - "return_log_probs": torch.zeros(metadata_len, dtype=torch.bool), - "top_n_logprobs": torch.zeros(metadata_len, dtype=torch.int64), - "termination_id": torch.full((metadata_len,), 99, dtype=torch.int64), + request_metadata=request_metadata, + active_request_metadata={ + label: metadata.clone() for label, metadata in request_metadata.items() }, + gpu_view=SimpleNamespace( + temperature=request_metadata["temperature"].clone(), + top_k=request_metadata["top_k"].to(torch.int32).clone(), + top_p=request_metadata["top_p"].clone(), + active_request_last_token_idxs=torch.arange(metadata_len, dtype=torch.int32), + request_query_lengths=torch.ones(metadata_len, dtype=torch.int32), + token_to_input_ids=torch.arange(32, dtype=torch.int64), + ), async_sched_step_count=0, async_sched_compaction_step_count=0, get_active_sequence_lengths=mock.Mock( @@ -225,9 +244,21 @@ def _make_async_sched_context(total_request_count=2, paused_request_count=0): return_value=torch.full((metadata_len,), 10, dtype=torch.int32) ), prepare_requests=mock.Mock(), - resolve_requests=mock.Mock(return_value=torch.empty(0, dtype=torch.int32)), + commit_sampled_tokens=mock.Mock(), + update_requests=mock.Mock(return_value={}), + resolve_requests=mock.Mock( + return_value=(torch.empty(0, dtype=torch.int32), torch.arange(metadata_len)) + ), + copy_async_sched_sample_to_forward=mock.Mock(), + reset=mock.Mock(), + transfer_bookkeeping_to_gpu=mock.Mock(return_value="bookkeeping"), using_cuda_graph_this_step=mock.Mock(return_value=False), + max_requests=metadata_len, + max_tokens=32, + request_query_lengths=torch.ones(metadata_len, dtype=torch.int32), ) + context.is_decode_only = mock.Mock(side_effect=lambda: context.num_prefill_requests == 0) + return context def _make_async_sched_controller(context=None, model_config=None): @@ -237,6 +268,7 @@ def _make_async_sched_controller(context=None, model_config=None): expert_model_parallel_size=1, num_moe_experts=None, moe_enable_routing_replay=False, + moe_pad_experts_for_cuda_graph_inference=False, ) controller = TextGenerationController.__new__(TextGenerationController) controller.inference_wrapped_model = SimpleNamespace( @@ -245,84 +277,197 @@ def _make_async_sched_controller(context=None, model_config=None): controller.model_config = model_config controller.num_speculative_tokens = 0 controller._enable_cuda_graph = False - controller._decode_forward_primer = DecodeForwardPrimer( - is_primed=True, cuda_graph_request_count=None + controller._sampling_backend = "torch" + controller._async_sched_logits = AsyncScheduleLogitsState(is_valid=True) + controller._async_sched_mtp_token_row_indices = None + controller._all_logits_cuda = torch.empty(0) + controller._sampled_tokens_cuda = torch.empty(context.max_requests, dtype=torch.int64) + + def sample_kernel(logits, n, _context, **kwargs): + sampled_tokens = torch.argmax(logits[:n], dim=-1) + output = kwargs.get("output") + if output is None: + return sampled_tokens + output.copy_(sampled_tokens) + return output + + controller._sampling = SimpleNamespace(sample_kernel=mock.Mock(side_effect=sample_kernel)) + controller._get_stop_word_finished_ids_callback = None + controller._async_sched_sampled_tokens_cpu_buffer = torch.empty( + context.max_requests, dtype=torch.int64 ) + controller._tp_size = 1 + controller._sp_enabled = False + controller._async_sched_selected_log_probs_cpu_buffer = torch.empty( + context.max_tokens, dtype=torch.float32 + ) + controller._async_sched_top_n_log_probs_cpu_buffer = None + controller._async_sched_top_n_token_ids_cpu_buffer = None + controller._async_sched_top_n_capacity = 0 return controller -def _set_nested_attr(obj, attr_path, value): - for attr in attr_path.split(".")[:-1]: - obj = getattr(obj, attr) - setattr(obj, attr_path.split(".")[-1], value) +@pytest.mark.parametrize( + "consumed, launched, expected", + [ + (False, False, False), + (True, True, True), + (None, None, ValueError), + (None, False, ValueError), + (None, True, ValueError), + (False, None, ValueError), + (True, None, ValueError), + (False, True, ValueError), + (True, False, ValueError), + ], +) +def test_decode_only_bool_requires_matching_forwards(consumed, launched, expected): + """Boolean conversion is valid only for matching consumed and launched forwards.""" + decode_only = DecodeOnly(consumed=consumed, launched=launched) + + if expected is ValueError: + with pytest.raises(ValueError, match="ambiguous"): + bool(decode_only) + else: + assert bool(decode_only) is expected + + +@pytest.mark.parametrize("num_prefill_requests, expected", [(1, False), (0, True)]) +def test_legacy_step_reports_matching_decode_only_state(num_prefill_requests, expected): + """Legacy consumes and launches the same computational batch.""" + context = _make_async_sched_context(total_request_count=1) + context.config.async_sched_mode = AsyncScheduleMode.LEGACY + context.num_prefill_requests = num_prefill_requests + context.num_decode_requests = 1 - num_prefill_requests + context.kv_block_allocator = SimpleNamespace(store_routing_per_block=mock.Mock()) + controller = _make_async_sched_controller(context) + controller._dynamic_step_context_init = mock.Mock( + return_value=(torch.tensor([1]), torch.tensor([0]), None) + ) + controller._dynamic_step_forward_logits = mock.Mock() + controller._router_record_bookkeeping = mock.Mock(return_value=None) + controller._dynamic_step_log_probs_bookkeeping = mock.Mock(return_value=(False, False)) + controller._dynamic_step_sample_logits = mock.Mock() + controller._dynamic_step_context_bookkeeping = mock.Mock( + return_value={"sample": torch.tensor([2])} + ) + + with mock.patch( + "megatron.core.inference.text_generation_controllers." + "text_generation_controller.get_moe_router_tracer", + return_value=None, + ): + result = asyncio.run(controller._run_legacy_step()) + + assert result.decode_only == DecodeOnly(consumed=expected, launched=expected) + assert result.output["sample"].tolist() == [2] @pytest.mark.parametrize("total_request_count", [0, 2]) def test_validate_async_sched_support_for_step_success(total_request_count): context = _make_async_sched_context(total_request_count=total_request_count) controller = _make_async_sched_controller(context) + if total_request_count == 0: + context.config.materialize_only_last_token_logits = False - controller._validate_async_sched_support_for_step() + controller._validate_async_sched_support_for_step(run_async_overlap=True) -@pytest.mark.parametrize( - "attr_path, value", - [ - ("context.config.materialize_only_last_token_logits", False), - ("controller.num_speculative_tokens", 1), - ("context.is_hybrid_model", True), - ("context.enable_prefix_caching", True), - ("context.paused_request_count", 1), - ("context.chunked_prefill_request_id", 0), - ("model_config.expert_model_parallel_size", 2), - ("model_config.num_moe_experts", 4), - ("model_config.moe_enable_routing_replay", True), - ("context.request_metadata", {"top_k": torch.tensor([1, 0])}), - ("context.request_metadata", {"top_p": torch.tensor([0.0, 0.5])}), - ("context.request_metadata", {"return_log_probs": torch.tensor([False, True])}), - ("context.request_metadata", {"top_n_logprobs": torch.tensor([0, 1])}), - ], -) -def test_validate_async_sched_support_for_step_errors(attr_path, value): +def test_validate_async_sched_support_for_step_allows_paused_no_overlap(): + """Paused requests are handled by no-overlap lifecycle bookkeeping.""" + context = _make_async_sched_context(total_request_count=2, paused_request_count=1) + controller = _make_async_sched_controller(context) + + controller._validate_async_sched_support_for_step(run_async_overlap=False) + + +def test_validate_async_sched_support_for_step_ignores_immutable_restrictions(): context = _make_async_sched_context(total_request_count=2) + context.config.materialize_only_last_token_logits = False + context.is_hybrid_model = True + context.enable_prefix_caching = True + context.chunked_prefill_request_id = 10 + context.request_metadata["top_k"] = torch.tensor([0, 0]) + context.request_metadata["top_p"] = torch.tensor([0.5, 0.5]) + context.request_metadata["return_log_probs"] = torch.tensor([True, True]) + context.request_metadata["top_n_logprobs"] = torch.tensor([1, 1]) model_config = SimpleNamespace( params_dtype=torch.float32, - expert_model_parallel_size=1, - num_moe_experts=None, - moe_enable_routing_replay=False, + expert_model_parallel_size=2, + num_moe_experts=4, + moe_enable_routing_replay=True, ) controller = _make_async_sched_controller(context, model_config) - target = SimpleNamespace(context=context, controller=controller, model_config=model_config) - if attr_path == "context.request_metadata": - context.request_metadata.update(value) - else: - _set_nested_attr(target, attr_path, value) + controller.num_speculative_tokens = 1 + + controller._validate_async_sched_support_for_step(run_async_overlap=True) + + +def test_validate_async_sched_support_for_step_errors_on_paused_overlap(): + context = _make_async_sched_context(total_request_count=2) + controller = _make_async_sched_controller(context) + context.paused_request_count = 1 with pytest.raises(RuntimeError, match="Async scheduling"): - controller._validate_async_sched_support_for_step() + controller._validate_async_sched_support_for_step(run_async_overlap=True) + + +def test_async_sched_logits_state_rejects_removed_ready_event(): + state = AsyncScheduleLogitsState() + + assert not hasattr(state, "ready_event") + with pytest.raises(TypeError): + AsyncScheduleLogitsState(ready_event=None) @pytest.mark.parametrize( - "enable_cuda_graph, survivor_idxs", + "enable_cuda_graph, survivor_idxs, expected_compaction", [ - (False, torch.tensor([0, 2], dtype=torch.int64)), - (True, torch.tensor([0, 2], dtype=torch.int64)), - (False, torch.empty(0, dtype=torch.int64)), + (False, torch.tensor([0, 2], dtype=torch.int64), True), + (True, torch.tensor([0, 2], dtype=torch.int64), True), + (False, torch.tensor([0, 1], dtype=torch.int64), False), + (False, torch.empty(0, dtype=torch.int64), False), ], ) -def test_async_sched_logits_compaction(enable_cuda_graph, survivor_idxs): - controller = _make_async_sched_controller() +def test_async_sched_logits_compaction(enable_cuda_graph, survivor_idxs, expected_compaction): + context = _make_async_sched_context(total_request_count=4) + context.active_request_metadata["temperature"].copy_(torch.tensor([0.1, 0.2, 0.3, 0.4])) + context.active_request_metadata["top_k"].copy_(torch.tensor([1, 2, 3, 4])) + context.active_request_metadata["top_p"].copy_(torch.tensor([0.5, 0.6, 0.7, 0.8])) + context.gpu_view.temperature.copy_(context.active_request_metadata["temperature"]) + context.gpu_view.top_k.copy_(context.active_request_metadata["top_k"]) + context.gpu_view.top_p.copy_(context.active_request_metadata["top_p"]) + original_cpu_metadata = { + label: context.active_request_metadata[label].clone() + for label in ("temperature", "top_k", "top_p") + } + gpu_metadata = { + "temperature": context.gpu_view.temperature, + "top_k": context.gpu_view.top_k, + "top_p": context.gpu_view.top_p, + } + original_gpu_metadata = {label: metadata.clone() for label, metadata in gpu_metadata.items()} + controller = _make_async_sched_controller(context) controller._enable_cuda_graph = enable_cuda_graph - controller._decode_forward_primer = DecodeForwardPrimer( - is_primed=True, cuda_graph_request_count=8 + controller._async_sched_logits = AsyncScheduleLogitsState( + is_valid=True, cuda_graph_request_count=8 ) logits = torch.arange(12).reshape(1, 4, 3) controller._all_logits_cuda = logits.clone() - controller._compact_async_sched_logits(survivor_idxs) + result = controller._compact_async_sched_logits(survivor_idxs) + + assert result is None if survivor_idxs.numel() == 0: - assert not controller._decode_forward_primer.is_primed + assert not controller._async_sched_logits.is_valid + return + + if not expected_compaction: + assert torch.equal(controller._all_logits_cuda, logits) + for label in ("temperature", "top_k", "top_p"): + assert torch.equal(context.active_request_metadata[label], original_cpu_metadata[label]) + assert torch.equal(gpu_metadata[label], original_gpu_metadata[label]) return expected_logits = logits[:, survivor_idxs, :] @@ -333,43 +478,122 @@ def test_async_sched_logits_compaction(enable_cuda_graph, survivor_idxs): assert controller._all_logits_cuda.shape == logits.shape else: assert torch.equal(controller._all_logits_cuda, expected_logits) - assert controller._decode_forward_primer.is_primed - assert controller._decode_forward_primer.cuda_graph_request_count == 8 + survivor_count = survivor_idxs.numel() + for label in ("temperature", "top_k", "top_p"): + assert torch.equal( + context.active_request_metadata[label][:survivor_count], + original_cpu_metadata[label][survivor_idxs], + ) + assert torch.equal( + gpu_metadata[label][:survivor_count], original_gpu_metadata[label][survivor_idxs] + ) + assert controller._async_sched_logits.is_valid + assert controller._async_sched_logits.cuda_graph_request_count == 8 + + +def test_async_sched_mtp_logits_compaction_preserves_input_rows(): + """MTP survivor logits retain the pending forward rows used for verification.""" + controller = _make_async_sched_controller(_make_async_sched_context(total_request_count=3)) + controller.num_speculative_tokens = 1 + controller._all_logits_cuda = torch.arange(18).reshape(1, 6, 3) + controller._async_sched_logits = AsyncScheduleLogitsState( + is_valid=True, token_row_indices=torch.tensor([10, 11, 20, 21, 30, 31]) + ) + + controller._compact_async_sched_logits(torch.tensor([2, 1])) + + assert controller._async_sched_logits.token_row_indices.tolist() == [30, 31, 20, 21] + + +def test_dynamic_step_context_init_returns_bookkeeping_event(): + context = _make_async_sched_context() + input_ids = torch.tensor([[10, 11]]) + position_ids = torch.tensor([[0, 1]]) + context.initialize_attention_state = mock.Mock(return_value="bookkeeping") + context.current_input_and_position_ids = mock.Mock(return_value=(input_ids, position_ids)) + model_config = SimpleNamespace( + params_dtype=torch.float32, + symmetric_ar_type=None, + nccl_all_reduce_for_prefill=False, + moe_pad_experts_for_cuda_graph_inference=False, + transformer_impl="transformer_engine", + ) + controller = _make_async_sched_controller(context, model_config) + + with ( + mock.patch( + "megatron.core.inference.text_generation_controllers." + "text_generation_controller.set_moe_metadata_sync" + ), + mock.patch( + "megatron.core.inference.text_generation_controllers.text_generation_controller.range_push" + ), + mock.patch( + "megatron.core.inference.text_generation_controllers.text_generation_controller.range_pop" + ), + ): + returned_input_ids, returned_position_ids, bookkeeping_done_event = ( + controller._dynamic_step_context_init(record_bookkeeping_done_event=True) + ) + + context.initialize_attention_state.assert_called_once_with( + construct_graph_dimensions=None, + is_expert_parallel_dummy_cuda_graph_step=False, + transfer_bookkeeping_to_gpu=True, + record_bookkeeping_done_event=True, + ) + assert torch.equal(returned_input_ids, input_ids) + assert torch.equal(returned_position_ids, position_ids) + assert bookkeeping_done_event == "bookkeeping" -def test_run_async_sched_prepare_updates_context_before_h2d_init(): +def test_run_async_sched_prepare_stops_before_bookkeeping_h2d(): context = _make_async_sched_context() controller = _make_async_sched_controller(context) input_ids = torch.tensor([[10, 11]]) position_ids = torch.tensor([[0, 1]]) call_order = [] - context.prepare_requests = mock.Mock(side_effect=lambda _: call_order.append("prepare")) + context.prepare_requests = mock.Mock(side_effect=lambda: call_order.append("prepare")) controller._dynamic_step_context_init = mock.Mock( - side_effect=lambda: call_order.append("context_init") or (input_ids, position_ids) + side_effect=lambda *, transfer_bookkeeping_to_gpu=True: call_order.append( + f"context_init:{transfer_bookkeeping_to_gpu}" + ) + or (input_ids, position_ids, None) ) - sample = torch.tensor([3, 4]) - returned_input_ids, returned_position_ids = controller._run_async_sched_prepare(sample) + returned_input_ids, returned_position_ids = controller._run_async_sched_prepare() - context.prepare_requests.assert_called_once_with(sample) + context.prepare_requests.assert_called_once_with() + controller._dynamic_step_context_init.assert_called_once_with(transfer_bookkeeping_to_gpu=False) assert torch.equal(returned_input_ids, input_ids) assert torch.equal(returned_position_ids, position_ids) - assert call_order == ["prepare", "context_init"] + assert call_order == ["prepare", "context_init:False"] + + +def test_run_async_sched_publish_bookkeeping_skips_gpu_input_ids(): + context = _make_async_sched_context() + controller = _make_async_sched_controller(context) + + done_event = controller._run_async_sched_publish_bookkeeping() + + context.transfer_bookkeeping_to_gpu.assert_called_once_with( + skip_token_input_ids=True, record_done_event=True + ) + assert done_event == "bookkeeping" @pytest.mark.parametrize( "using_cuda_graph, expected_cuda_graph_request_count", [(False, None), (True, 8)] ) -def test_run_async_sched_forward_records_primer( +def test_run_async_sched_forward_records_pending_logits( using_cuda_graph, expected_cuda_graph_request_count ): context = _make_async_sched_context() context.using_cuda_graph_this_step.return_value = using_cuda_graph controller = _make_async_sched_controller(context) - controller._decode_forward_primer = DecodeForwardPrimer( - is_primed=False, cuda_graph_request_count=None - ) + controller._async_sched_logits = AsyncScheduleLogitsState() + controller._all_logits_cuda = torch.empty(0) controller._dynamic_step_forward_logits = mock.Mock() input_ids = torch.tensor([[10, 11]]) position_ids = torch.tensor([[0, 1]]) @@ -384,126 +608,1026 @@ def test_run_async_sched_forward_records_primer( "text_generation_controller.range_pop" ), ): - cuda_graph_request_count = controller._run_async_sched_forward(input_ids, position_ids) + result = controller._run_async_sched_forward(input_ids, position_ids) controller._dynamic_step_forward_logits.assert_called_once_with(input_ids, position_ids) - assert cuda_graph_request_count == expected_cuda_graph_request_count - assert controller._decode_forward_primer.is_primed + assert result is None + assert controller._async_sched_logits.is_valid assert ( - controller._decode_forward_primer.cuda_graph_request_count - == expected_cuda_graph_request_count + controller._async_sched_logits.cuda_graph_request_count == expected_cuda_graph_request_count + ) + + +def test_run_async_sched_forward_commits_mamba_prefix_states(): + """A real async forward commits cacheable Mamba states before row resolution.""" + context = _make_async_sched_context() + context.is_hybrid_model = True + context.mamba_slot_allocator = SimpleNamespace(commit_intermediate_states=mock.Mock()) + controller = _make_async_sched_controller(context) + controller._dynamic_step_forward_logits = mock.Mock() + + controller._run_async_sched_forward(torch.tensor([[10, 11]]), torch.tensor([[0, 1]])) + + context.mamba_slot_allocator.commit_intermediate_states.assert_called_once_with() + + +def test_run_dummy_async_sched_base_step_resets_without_committing_mamba_state(): + context = _make_async_sched_context() + context.is_hybrid_model = True + context.mamba_slot_allocator = SimpleNamespace(commit_intermediate_states=mock.Mock()) + controller = _make_async_sched_controller(context) + input_ids = torch.tensor([[10]]) + position_ids = torch.tensor([[0]]) + controller._dynamic_step_context_init = mock.Mock(return_value=(input_ids, position_ids, None)) + controller._run_dummy_base_forward = mock.Mock() + + result = controller._run_dummy_async_sched_base_step() + + assert result is None + controller._run_dummy_base_forward.assert_called_once_with(input_ids, position_ids) + context.mamba_slot_allocator.commit_intermediate_states.assert_not_called() + context.reset.assert_called_once_with(preserve_prefix_cache=True, preserve_counters=True) + + +@pytest.mark.parametrize("is_valid", [False, True]) +def test_run_async_sched_forward_primer(is_valid): + context = _make_async_sched_context(total_request_count=2) + controller = _make_async_sched_controller(context) + controller._async_sched_logits = AsyncScheduleLogitsState(is_valid=is_valid) + input_ids = torch.tensor([[10, 11]]) + position_ids = torch.tensor([[0, 1]]) + controller._dynamic_step_context_init = mock.Mock( + return_value=(input_ids, position_ids, "bookkeeping") + ) + controller._run_async_sched_forward = mock.Mock() + + primer_launched, bookkeeping_done_event = controller._run_async_sched_forward_primer() + + assert primer_launched is (not is_valid) + assert bookkeeping_done_event == (None if is_valid else "bookkeeping") + if is_valid: + controller._dynamic_step_context_init.assert_not_called() + controller._run_async_sched_forward.assert_not_called() + else: + controller._dynamic_step_context_init.assert_called_once_with( + record_bookkeeping_done_event=True + ) + controller._run_async_sched_forward.assert_called_once_with(input_ids, position_ids) + + +@pytest.mark.parametrize( + "mode, expected_order", + [(AsyncScheduleMode.LEGACY, ["base", "mtp"]), (AsyncScheduleMode.ASYNC, ["mtp", "base"])], +) +def test_dummy_forward_matches_mode_forward_order(mode, expected_order): + """Real and idle EP ranks issue MTP and base collectives in the same order.""" + context = _make_async_sched_context() + context.config.async_sched_mode = mode + controller = _make_async_sched_controller(context) + input_ids = torch.tensor([[1]]) + position_ids = torch.tensor([[0]]) + order = [] + controller._dynamic_step_context_init = mock.Mock(return_value=(input_ids, position_ids, None)) + controller._run_dummy_base_forward = mock.Mock(side_effect=lambda *_: order.append("base")) + controller._run_dummy_serial_mtp_forward = mock.Mock(side_effect=lambda: order.append("mtp")) + + controller.dummy_forward() + + assert order == expected_order + context.reset.assert_called_once_with(preserve_prefix_cache=True, preserve_counters=True) + + +@pytest.mark.parametrize( + "num_speculative_tokens, ep_size, expected_order", + [(0, 2, ["base"]), (2, 1, ["base"]), (2, 2, ["mtp", "base"])], +) +def test_async_sched_primer_matches_dummy_mtp_order( + num_speculative_tokens, ep_size, expected_order +): + """An EP primer inserts dummy MTP collectives only when its peers do.""" + context = _make_async_sched_context(total_request_count=1) + model_config = SimpleNamespace( + params_dtype=torch.float32, + expert_model_parallel_size=ep_size, + num_moe_experts=4 if ep_size > 1 else None, + moe_enable_routing_replay=False, ) + controller = _make_async_sched_controller(context, model_config) + controller.num_speculative_tokens = num_speculative_tokens + controller._async_sched_logits = AsyncScheduleLogitsState() + controller._dynamic_step_context_init = mock.Mock( + return_value=(torch.tensor([[1]]), torch.tensor([[0]]), "bookkeeping") + ) + order = [] + controller._run_dummy_serial_mtp_forward = mock.Mock(side_effect=lambda: order.append("mtp")) + controller._run_async_sched_forward = mock.Mock(side_effect=lambda *_: order.append("base")) + + controller._run_async_sched_forward_primer() + assert order == expected_order -def test_async_sched_serial_step_returns_none_without_active_requests(): + +def test_async_sched_router_returns_empty_result_without_active_requests(): context = _make_async_sched_context(total_request_count=0) context.active_token_count = 0 controller = _make_async_sched_controller(context) - controller._decode_forward_primer = DecodeForwardPrimer( - is_primed=True, cuda_graph_request_count=8 + controller._async_sched_logits = AsyncScheduleLogitsState( + is_valid=True, cuda_graph_request_count=8 ) controller._validate_async_sched_support_for_step = mock.Mock() - result = asyncio.run(controller._run_async_sched_serial_step()) + result = asyncio.run(controller.async_generate_output_tokens_dynamic_batch()) - assert result is None - assert not controller._decode_forward_primer.is_primed - controller._validate_async_sched_support_for_step.assert_not_called() + assert result == DynamicBatchControllerStepResult( + decode_only=DecodeOnly(consumed=None, launched=None) + ) + assert not controller._async_sched_logits.is_valid + assert controller._async_sched_logits.cuda_graph_request_count is None + controller._validate_async_sched_support_for_step.assert_called_once_with(True) + + +@pytest.mark.parametrize("logits_dtype", [torch.float32, torch.bfloat16]) +def test_run_async_sched_sample_reuses_gpu_buffer(logits_dtype): + context = _make_async_sched_context(total_request_count=3) + controller = _make_async_sched_controller(context) + controller._all_logits_cuda = torch.zeros(1, 3, 5, dtype=logits_dtype) + expected_tokens = torch.tensor([1, 2, 3], dtype=torch.int64) + for idx, token in enumerate(expected_tokens.tolist()): + controller._all_logits_cuda[0, idx, token] = 10.0 + + result = controller._run_async_sched_sample() + + assert result.sampled_tokens_gpu.data_ptr() == controller._sampled_tokens_cuda.data_ptr() + assert torch.equal(result.sampled_tokens_gpu, expected_tokens) + assert torch.equal(result.sampled_tokens_cpu_view, expected_tokens) + controller._sampling.sample_kernel.assert_called_once() + logits, n, called_context = controller._sampling.sample_kernel.call_args.args + assert torch.equal(logits, controller._all_logits_cuda.squeeze(0)) + assert n == 3 + assert called_context is context + sample_kwargs = controller._sampling.sample_kernel.call_args.kwargs + assert set(sample_kwargs) == {"gather_indices", "no_top_k", "no_top_p", "output"} + assert sample_kwargs["gather_indices"] is None + assert not sample_kwargs["no_top_k"] + assert sample_kwargs["no_top_p"] + assert sample_kwargs["output"].data_ptr() == controller._sampled_tokens_cuda.data_ptr() + + +def test_async_sched_log_probs_materializes_decode_top_n_without_padding(): + """Async logprobs exclude padded graph rows from top-n results.""" + context = _make_async_sched_context(total_request_count=3) + context.num_decode_requests = 3 + context.active_request_metadata["return_log_probs"].fill_(True) + context.active_request_metadata["top_n_logprobs"].copy_(torch.tensor([0, 2, 1])) + context.calculate_log_probs_tensors = mock.Mock( + return_value=( + torch.tensor([-0.1, -0.2, -0.3]), + torch.tensor( + [ + [-4.0, -1.0, -3.0, -2.0], + [-0.4, -0.1, -0.3, -0.2], + [-2.0, -4.0, -1.0, -3.0], + [10.0, 9.0, 8.0, 7.0], + ] + ), + ) + ) + controller = _make_async_sched_controller(context) + controller._all_logits_cuda = torch.empty(1, 3, 4) + sampled_tokens = torch.tensor([1, 1, 2]) + + gpu_result = controller._run_async_sched_log_probs( + SimpleNamespace(sampled_tokens_gpu=sampled_tokens, accepted_counts_gpu=None) + ) + transfer = controller._copy_async_sched_log_probs_to_cpu(gpu_result) + log_probs, top_n_logprobs = controller._materialize_async_sched_log_probs(transfer) + + assert [values[0] for values in log_probs] == pytest.approx([-0.1, -0.2, -0.3]) + assert set(top_n_logprobs) == {1, 2} + assert torch.equal(top_n_logprobs[1][0][1], torch.tensor([1, 3])) + assert torch.equal(top_n_logprobs[2][0][1], torch.tensor([2])) + + +def test_async_sched_log_probs_materializes_mtp_and_prefill_rows(): + """Async logprobs preserve variable MTP acceptance and prompt row boundaries.""" + context = _make_async_sched_context(total_request_count=3) + context.config.materialize_only_last_token_logits = False + context.num_decode_requests = 2 + context.num_prefill_requests = 1 + context.active_token_count = 9 + context.request_query_lengths.copy_(torch.tensor([3, 3, 3])) + context.gpu_view.request_query_lengths.copy_(torch.tensor([3, 3, 3])) + context.gpu_view.token_to_input_ids[:9].copy_(torch.tensor([1, 2, 3, 4, 5, 6, 30, 31, 32])) + context.active_request_metadata["return_log_probs"].fill_(True) + context.active_request_metadata["top_n_logprobs"].copy_(torch.tensor([0, 0, 2])) + context.active_request_metadata["skip_prompt_log_probs"].copy_( + torch.tensor([False, False, True]) + ) + context.calculate_log_probs_tensors = mock.Mock( + return_value=( + -torch.arange(1, 10, dtype=torch.float32) / 10, + torch.arange(36, dtype=torch.float32).view(9, 4), + ) + ) + controller = _make_async_sched_controller(context) + controller.num_speculative_tokens = 2 + controller._all_logits_cuda = torch.empty(1, 9, 4) + controller._accepted_tokens_per_request = torch.tensor([[10, -1], [11, 12], [-1, -1]]) + sample_result = SimpleNamespace( + sampled_tokens_gpu=torch.tensor([20, 21, 22]), accepted_counts_gpu=torch.tensor([1, 2, 0]) + ) + + gpu_result = controller._run_async_sched_log_probs(sample_result) + transfer = controller._copy_async_sched_log_probs_to_cpu(gpu_result) + log_probs, top_n_logprobs = controller._materialize_async_sched_log_probs( + transfer, torch.tensor([1, 2, 0]) + ) + + assert log_probs[0] == pytest.approx([-0.1, -0.2]) + assert log_probs[1] == pytest.approx([-0.4, -0.5, -0.6]) + assert log_probs[2] == pytest.approx([-0.7, -0.8, -0.9]) + assert len(top_n_logprobs[2]) == 1 + assert torch.equal(top_n_logprobs[2][0][1], torch.tensor([3, 2])) + assert torch.equal( + context.calculate_log_probs_tensors.call_args.args[1], + torch.tensor([10, 20, 20, 11, 12, 21, 31, 32, 22]), + ) + assert torch.equal( + context.calculate_log_probs_tensors.call_args.kwargs["row_to_request"], + torch.tensor([0, 0, 0, 1, 1, 1, 2, 2, 2]), + ) + + +@pytest.mark.internal +def test_run_async_sched_sample_records_gpu_ready_event(): + context = _make_async_sched_context(total_request_count=3) + controller = _make_async_sched_controller(context) + controller._all_logits_cuda = torch.zeros(1, 3, 5, device="cuda") + controller._sampled_tokens_cuda = torch.empty(3, dtype=torch.int64, device="cuda") + controller._async_sched_sample_gpu_ready_event = mock.Mock() + controller._copy_async_sched_sample_to_cpu = mock.Mock( + return_value=(torch.empty(3), None, None, "sample_cpu") + ) + + controller._run_async_sched_sample() + + controller._async_sched_sample_gpu_ready_event.record.assert_called_once_with( + torch.cuda.current_stream() + ) + + +@pytest.mark.internal +def test_synchronize_async_sched_event_handles_cuda_event_and_none(): + controller = _make_async_sched_controller() + event = torch.cuda.Event() + event.record(torch.cuda.current_stream()) + controller._synchronize_async_sched_event(event) + controller._synchronize_async_sched_event(None) + + assert isinstance(event, torch.cuda.Event) + + +@pytest.mark.internal +def test_copy_async_sched_sample_to_cpu_uses_reusable_buffer_and_copy_stream(): + controller = _make_async_sched_controller(_make_async_sched_context(total_request_count=3)) + sampled_tokens_gpu = torch.tensor([1, 2, 3], dtype=torch.int64, device="cuda") + controller._async_sched_sampled_tokens_cpu_buffer = torch.empty( + 3, dtype=torch.int64, device="cpu", pin_memory=True + ) + controller._async_sched_sample_gpu_ready_event = torch.cuda.Event() + controller._async_sched_sample_cpu_ready_event = torch.cuda.Event() + controller._async_sched_copy_stream = torch.cuda.Stream() + controller._async_sched_sample_gpu_ready_event.record(torch.cuda.current_stream()) + + sampled_tokens_cpu_view, sampled_mtp_tokens_cpu, accepted_tokens_cpu, sample_cpu_ready_event = ( + controller._copy_async_sched_sample_to_cpu(sampled_tokens_gpu) + ) + sample_cpu_ready_event.synchronize() + + assert ( + sampled_tokens_cpu_view.data_ptr() + == controller._async_sched_sampled_tokens_cpu_buffer.data_ptr() + ) + assert torch.equal(sampled_tokens_cpu_view, sampled_tokens_gpu.cpu()) + assert sampled_mtp_tokens_cpu is None + assert accepted_tokens_cpu is None + + +@pytest.mark.internal +def test_copy_async_sched_log_probs_to_cpu_uses_reusable_buffers_and_copy_stream(): + """Async logprob D2H copies reuse pinned storage and report completion.""" + context = _make_async_sched_context(total_request_count=3) + context.num_decode_requests = 3 + context.active_request_metadata["return_log_probs"].fill_(True) + context.active_request_metadata["top_n_logprobs"].copy_(torch.tensor([0, 2, 1])) + context.calculate_log_probs_tensors = mock.Mock( + return_value=( + torch.tensor([-0.1, -0.2, -0.3], device="cuda"), + torch.tensor( + [[-4.0, -1.0, -3.0, -2.0], [-0.4, -0.1, -0.3, -0.2], [-2.0, -4.0, -1.0, -3.0]], + device="cuda", + ), + ) + ) + controller = _make_async_sched_controller(context) + controller._all_logits_cuda = torch.empty(1, 3, 4, device="cuda") + controller._async_sched_selected_log_probs_cpu_buffer = torch.empty( + 3, dtype=torch.float32, device="cpu", pin_memory=True + ) + controller._async_sched_log_probs_gpu_ready_event = torch.cuda.Event() + controller._async_sched_log_probs_cpu_ready_event = torch.cuda.Event() + controller._async_sched_copy_stream = torch.cuda.Stream() + + gpu_result = controller._run_async_sched_log_probs( + SimpleNamespace( + sampled_tokens_gpu=torch.tensor([1, 1, 2], device="cuda"), accepted_counts_gpu=None + ) + ) + transfer = controller._copy_async_sched_log_probs_to_cpu(gpu_result) + transfer.cpu_ready_event.synchronize() + log_probs, top_n_logprobs = controller._materialize_async_sched_log_probs(transfer) + + assert [values[0] for values in log_probs] == pytest.approx([-0.1, -0.2, -0.3]) + assert torch.equal(top_n_logprobs[1][0][1], torch.tensor([1, 3])) + assert torch.equal(top_n_logprobs[2][0][1], torch.tensor([2])) + + +def test_build_async_sched_request_state_uses_resolved_lengths(): + """Resolution tests the accepted output length, not speculative prepared state.""" + context = _make_async_sched_context(total_request_count=2) + context.get_max_sequence_lengths.return_value = torch.tensor([4, 4]) + controller = _make_async_sched_controller(context) + + _, finished_request_ids, active_mask = controller._build_async_sched_request_state( + torch.tensor([1, 2]), torch.tensor([3, 4]) + ) + + assert active_mask.tolist() == [1, 0] + assert finished_request_ids.tolist() == [11] + + +def test_build_async_sched_request_state_keeps_partial_chunk_active(): + """A partial chunk ignores its provisional sample and remains schedulable.""" + context = _make_async_sched_context(total_request_count=3) + context.chunked_prefill_request_id = 12 + context.get_max_sequence_lengths.return_value = torch.tensor([4, 4, 4]) + controller = _make_async_sched_controller(context) + + _, finished_request_ids, active_mask = controller._build_async_sched_request_state( + torch.tensor([1, 2, 3]), torch.tensor([3, 4, 4]) + ) + + assert active_mask.tolist() == [1, 0, 1] + assert finished_request_ids.tolist() == [11] @pytest.mark.parametrize( - "is_primed, termination_ids, expected_finished_ids, expected_compaction_count", - [(True, torch.tensor([99, 99, 99]), [], 0), (False, torch.tensor([99, 2, 99]), [11], 1)], + "termination_ids, stop_word_finished_ids", + [([99, 99, 99], set()), ([99, 2, 99], set()), ([99, 99, 99], {11})], ) -def test_async_sched_serial_step( - is_primed, termination_ids, expected_finished_ids, expected_compaction_count +def test_run_async_sched_resolve_compacts_without_forward_sync( + termination_ids, stop_word_finished_ids ): sample_tokens = torch.tensor([1, 2, 3], dtype=torch.int64) context = _make_async_sched_context(total_request_count=3) - context.request_metadata["termination_id"] = termination_ids + context.request_metadata["termination_id"] = torch.tensor(termination_ids) + controller = _make_async_sched_controller(context) + controller._synchronize_async_sched_event = mock.Mock() + controller._get_stop_word_finished_ids_callback = mock.Mock(return_value=stop_word_finished_ids) + + expected_mask = (sample_tokens != context.request_metadata["termination_id"]).byte() + for request_idx, request_id in enumerate(context.request_ids.tolist()): + if request_id in stop_word_finished_ids: + expected_mask[request_idx] = 0 + expected_finished_ids = context.request_ids[expected_mask == 0].clone() + expected_survivor_idxs = torch.nonzero(expected_mask, as_tuple=True)[0] context.resolve_requests = mock.Mock( - return_value=torch.tensor(expected_finished_ids, dtype=torch.int32) + return_value=(expected_finished_ids, expected_survivor_idxs) ) - controller = _make_async_sched_controller(context) - controller._decode_forward_primer = DecodeForwardPrimer( - is_primed=is_primed, cuda_graph_request_count=7 if is_primed else None + + controller._compact_async_sched_logits = mock.Mock() + + sample_result = SimpleNamespace( + sampled_tokens_cpu_view=sample_tokens, accepted_tokens_cpu_view=None ) - controller._validate_async_sched_support_for_step = mock.Mock() - controller._all_logits_cuda = torch.zeros(1, 3, 5) - for idx, token in enumerate(sample_tokens.tolist()): - controller._all_logits_cuda[0, idx, token] = 10.0 + result = controller._run_async_sched_resolve( + sample_result, context.get_active_sequence_lengths() + 1 + ) + + assert torch.equal(result.sampled_tokens_cpu, sample_tokens) + assert not hasattr(result, "compaction_done_event") + controller._synchronize_async_sched_event.assert_not_called() + assert torch.equal(result.survivor_idxs, expected_survivor_idxs) + controller._compact_async_sched_logits.assert_called_once_with(expected_survivor_idxs) + context.commit_sampled_tokens.assert_not_called() + context.resolve_requests.assert_called_once() + assert torch.equal(context.resolve_requests.call_args.args[0], expected_mask) + controller._get_stop_word_finished_ids_callback.assert_called_once_with([10, 11, 12]) + +def test_async_sched_step_overlap_order(): + """Logprob transfer overlaps after current-logit GPU work is queued.""" + sample_tokens = torch.tensor([1, 2, 3], dtype=torch.int64) + sampled_tokens_cpu = sample_tokens.clone() input_ids = torch.tensor([[101, 102, 103]]) position_ids = torch.tensor([[0, 1, 2]]) + context = _make_async_sched_context(total_request_count=3) + controller = _make_async_sched_controller(context) + controller._async_sched_logits = AsyncScheduleLogitsState( + is_valid=True, cuda_graph_request_count=7 + ) call_order = [] - controller._dynamic_step_context_init = mock.Mock( - side_effect=lambda: call_order.append("context_init") or (input_ids, position_ids) + + controller._synchronize_async_sched_event = mock.Mock( + side_effect=lambda event: call_order.append(f"wait:{event}") + ) + controller._run_async_sched_prepare = mock.Mock( + side_effect=lambda: call_order.append("prepare") or (input_ids, position_ids) + ) + controller._run_async_sched_sample = mock.Mock( + side_effect=lambda: call_order.append("sample") + or SimpleNamespace( + sampled_tokens_gpu=sample_tokens, + sampled_tokens_cpu_view=sampled_tokens_cpu, + sampled_mtp_tokens_gpu=None, + sampled_mtp_tokens_cpu_view=None, + accepted_tokens_cpu_view=None, + sample_cpu_ready_event="sample", + ) + ) + context.copy_async_sched_sample_to_forward = mock.Mock( + side_effect=lambda _: call_order.append("copy_input") + ) + log_probs_gpu_result = object() + log_probs_transfer = SimpleNamespace(cpu_ready_event="log_probs") + controller._run_async_sched_log_probs = mock.Mock( + side_effect=lambda _: call_order.append("log_probs") or log_probs_gpu_result + ) + controller._copy_async_sched_log_probs_to_cpu = mock.Mock( + side_effect=lambda _: call_order.append("copy_log_probs") or log_probs_transfer + ) + controller._materialize_async_sched_log_probs = mock.Mock( + side_effect=lambda *_: call_order.append("materialize_log_probs") + or ([[0.1], [0.2], [0.3]], None) + ) + context.commit_sampled_tokens = mock.Mock(side_effect=lambda *_: call_order.append("commit")) + controller._run_async_sched_publish_bookkeeping = mock.Mock( + side_effect=lambda: call_order.append("publish") or "bookkeeping" + ) + controller._run_async_sched_forward = mock.Mock( + side_effect=lambda *_: call_order.append("forward") + ) + controller._run_async_sched_resolve = mock.Mock( + side_effect=lambda *_: call_order.append("resolve") + or SimpleNamespace( + sampled_tokens_cpu=sampled_tokens_cpu, + accepted_tokens_cpu=None, + active_request_ids=context.request_ids.long(), + finished_request_ids=torch.tensor([11], dtype=torch.int32), + survivor_idxs=torch.tensor([0, 2]), + newly_paused_request_ids=None, + evict_request_ids=None, + ) ) - def forward_step(forward_input_ids, forward_position_ids): - call_order.append("forward") - assert torch.equal(forward_input_ids, input_ids) - assert torch.equal(forward_position_ids, position_ids) - controller._decode_forward_primer.mark_primed(5) - return 5 + async def yield_to_event_loop(_delay): + call_order.append("yield") - controller._run_async_sched_forward = mock.Mock(side_effect=forward_step) - context.prepare_requests = mock.Mock(side_effect=lambda _: call_order.append("prepare")) + with mock.patch( + "megatron.core.inference.text_generation_controllers." + "text_generation_controller.asyncio.sleep", + side_effect=yield_to_event_loop, + ): + result = asyncio.run(controller._run_async_sched_step_overlap()) - def compact_logits(survivor_idxs): - assert context.async_sched_step_count == 0 - assert context.async_sched_compaction_step_count == 0 - expected_survivors = torch.tensor( - [idx for idx, token in enumerate(sample_tokens.tolist()) if token != 2], - dtype=torch.int64, - ) - if not expected_finished_ids: - expected_survivors = torch.arange(sample_tokens.numel(), dtype=torch.int64) - assert torch.equal(survivor_idxs, expected_survivors) + assert result.output["sample"].tolist() == sample_tokens.tolist() + assert result.decode_only == DecodeOnly(consumed=True, launched=True) + assert result.output["log_probs"] == [[0.1], [0.2], [0.3]] + assert result.output["cuda_graph_request_count"] == 7 + assert context.async_sched_step_count == 1 + assert context.async_sched_compaction_step_count == 1 + assert call_order == [ + "prepare", + "sample", + "copy_input", + "log_probs", + "copy_log_probs", + "publish", + "forward", + "wait:sample", + "wait:bookkeeping", + "resolve", + "commit", + "wait:log_probs", + "materialize_log_probs", + "yield", + ] + assert torch.equal( + context.commit_sampled_tokens.call_args.args[0], sampled_tokens_cpu[torch.tensor([0, 2])] + ) - controller._compact_async_sched_logits = mock.Mock(side_effect=compact_logits) - result = asyncio.run(controller._run_async_sched_serial_step()) +@pytest.mark.parametrize( + "termination_ids, expected_mask, expected_finished_ids, expected_survivor_idxs, " + "expected_compaction_count", + [ + ([99, 99, 99], [1, 1, 1], [], [0, 1, 2], 0), + ([99, 2, 99], [1, 0, 1], [11], [0, 2], 1), + ([1, 99, 99], [0, 1, 1], [10], [2, 1], 1), + ([1, 2, 3], [0, 0, 0], [10, 11, 12], [], 1), + ], +) +def test_async_sched_step_wires_sampling_through_resolution( + termination_ids, + expected_mask, + expected_finished_ids, + expected_survivor_idxs, + expected_compaction_count, +): + context = _make_async_sched_context(total_request_count=3) + context.request_metadata["termination_id"] = torch.tensor(termination_ids) + context.resolve_requests.side_effect = lambda mask: ( + context.request_ids[mask == 0].clone(), + torch.tensor(expected_survivor_idxs, dtype=torch.long), + ) + controller = _make_async_sched_controller(context) + controller._all_logits_cuda = torch.zeros(1, 3, 5) + sampled_tokens = torch.tensor([1, 2, 3], dtype=torch.int64) + for row, token in enumerate(sampled_tokens.tolist()): + controller._all_logits_cuda[0, row, token] = 10.0 + + input_ids_gpu_view = torch.empty(3, dtype=torch.int64) + position_ids_gpu_view = torch.empty(3, dtype=torch.int64) + controller._run_async_sched_prepare = mock.Mock( + return_value=(input_ids_gpu_view, position_ids_gpu_view) + ) + controller._run_async_sched_publish_bookkeeping = mock.Mock(return_value=None) + controller._synchronize_async_sched_event = mock.Mock() + + def run_forward(*_args): + controller._async_sched_logits.set_pending(None) + + controller._run_async_sched_forward = mock.Mock(side_effect=run_forward) + step_result = asyncio.run(controller._run_async_sched_step_overlap()) + result = step_result.output + + assert torch.equal(result["sample"], sampled_tokens) assert result["finished_request_ids"].tolist() == expected_finished_ids - assert result["sample"].tolist() == sample_tokens.tolist() - assert result["cuda_graph_request_count"] == (7 if is_primed else 5) + context.copy_async_sched_sample_to_forward.assert_called_once() + assert torch.equal(context.copy_async_sched_sample_to_forward.call_args.args[0], sampled_tokens) + context.commit_sampled_tokens.assert_called_once() + assert torch.equal( + context.commit_sampled_tokens.call_args.args[0], + sampled_tokens[torch.tensor(expected_survivor_idxs, dtype=torch.long)], + ) + assert context.resolve_requests.call_args.args[0].tolist() == expected_mask assert context.async_sched_step_count == 1 assert context.async_sched_compaction_step_count == expected_compaction_count - context.prepare_requests.assert_called_once() - context.resolve_requests.assert_called_once() - controller._compact_async_sched_logits.assert_called_once() - expected_prefix = [] if is_primed else ["context_init", "forward"] - assert call_order == expected_prefix + ["prepare", "context_init", "forward"] + + +def test_async_sched_step_yields_after_resolution_outside_inference_mode(): + context = _make_async_sched_context(total_request_count=1) + controller = _make_async_sched_controller(context) + sampled_tokens = torch.tensor([1], dtype=torch.int64) + controller._run_async_sched_prepare = mock.Mock( + return_value=(torch.empty(1, dtype=torch.int64), torch.empty(1, dtype=torch.int64)) + ) + controller._run_async_sched_sample = mock.Mock( + return_value=SimpleNamespace( + sampled_tokens_gpu=sampled_tokens, + sampled_tokens_cpu_view=sampled_tokens, + sampled_mtp_tokens_gpu=None, + sampled_mtp_tokens_cpu_view=None, + accepted_tokens_cpu_view=None, + sample_cpu_ready_event=None, + ) + ) + controller._run_async_sched_publish_bookkeeping = mock.Mock(return_value=None) + controller._run_async_sched_forward = mock.Mock(return_value=None) + controller._run_async_sched_resolve = mock.Mock( + return_value=SimpleNamespace( + sampled_tokens_cpu=sampled_tokens, + accepted_tokens_cpu=None, + active_request_ids=context.request_ids.long(), + finished_request_ids=torch.empty(0, dtype=torch.int32), + survivor_idxs=torch.tensor([0]), + newly_paused_request_ids=None, + evict_request_ids=None, + ) + ) + observed = [] + + async def run_step(): + asyncio.get_running_loop().call_soon( + lambda: observed.append( + (context.async_sched_step_count, torch.is_inference_mode_enabled()) + ) + ) + return await controller._run_async_sched_step_overlap() + + result = asyncio.run(run_step()).output + + assert result["sample"].tolist() == [1] + assert observed == [(1, False)] + + +def test_async_sched_initial_no_overlap_step_launches_primer_only(): + """Initial admission launches one primer and returns across the engine boundary.""" + context = _make_async_sched_context(total_request_count=0) + context.active_token_count = 0 + controller = _make_async_sched_controller(context) + controller._async_sched_logits = AsyncScheduleLogitsState() + call_order = [] + + def admit_request(): + call_order.append("admit") + context.total_request_count = 1 + context.active_token_count = 4 + context.num_prefill_requests = 1 + + controller._run_async_sched_forward_primer = mock.Mock( + side_effect=lambda: call_order.append("primer") or (True, "bookkeeping") + ) + controller._synchronize_async_sched_event = mock.Mock( + side_effect=lambda event: call_order.append(f"wait:{event}") + ) + + result = asyncio.run( + controller._run_async_sched_step_no_overlap(schedule_waiting_requests=admit_request) + ) + + assert result == DynamicBatchControllerStepResult( + decode_only=DecodeOnly(consumed=None, launched=False), primer_only=True + ) + assert call_order == ["admit", "primer", "wait:bookkeeping"] + assert context.async_sched_step_count == 0 + + +def test_run_async_sched_update_requests_preserves_pre_update_output(): + """Legacy bookkeeping may reorder working samples without corrupting step output.""" + context = _make_async_sched_context(total_request_count=3, paused_request_count=1) + context.request_metadata["termination_id"] = torch.tensor([99, 99, 2]) + context.get_max_sequence_lengths.return_value = torch.tensor([10, 10]) + controller = _make_async_sched_controller(context) + sampled_tokens = torch.tensor([1, 2]) + sampled_mtp_tokens = torch.tensor([[3, 4], [5, 6]]) + accepted_tokens = torch.tensor([7, 8]) + sample_result = SimpleNamespace( + sampled_tokens_cpu_view=sampled_tokens, + sampled_mtp_tokens_cpu_view=sampled_mtp_tokens, + accepted_tokens_cpu_view=accepted_tokens, + ) + + def update_requests(active_mask, mutable_samples, mutable_mtp_samples): + assert active_mask.tolist() == [1, 0] + mutable_samples.fill_(-1) + mutable_mtp_samples.fill_(-1) + return { + "newly_paused_request_ids": torch.tensor([11]), + "evict_request_ids": torch.tensor([10]), + } + + context.update_requests.side_effect = update_requests + + result = controller._run_async_sched_update_requests( + sample_result, resolved_sequence_lengths=torch.tensor([4, 4]) + ) + + assert result.active_request_ids.tolist() == [11, 12] + assert result.finished_request_ids.tolist() == [12] + assert result.sampled_tokens_cpu.tolist() == [1, 2] + assert result.accepted_tokens_cpu.tolist() == [7, 8] + assert result.newly_paused_request_ids.tolist() == [11] + assert result.evict_request_ids.tolist() == [10] + assert sampled_tokens.tolist() == [1, 2] + assert sampled_mtp_tokens.tolist() == [[3, 4], [5, 6]] @pytest.mark.parametrize( - "mode, num_prefill_requests, skip_bookkeeping, expected_result", + "consumed_prefill_requests, launched_prefill_requests", [(1, 0), (1, 1), (0, 0), (0, 1)] +) +def test_async_sched_no_overlap_updates_before_admission( + consumed_prefill_requests, launched_prefill_requests +): + """No-overlap classifies consumed and launched work around lifecycle updates.""" + context = _make_async_sched_context(total_request_count=2) + context.num_prefill_requests = consumed_prefill_requests + controller = _make_async_sched_controller(context) + controller._async_sched_logits = AsyncScheduleLogitsState( + is_valid=True, cuda_graph_request_count=7 + ) + sampled_tokens = torch.tensor([1, 2]) + sample_result = SimpleNamespace( + sampled_tokens_gpu=sampled_tokens, + sampled_tokens_cpu_view=sampled_tokens, + sampled_mtp_tokens_gpu=None, + sampled_mtp_tokens_cpu_view=None, + accepted_tokens_cpu_view=None, + sample_cpu_ready_event="sample", + ) + request_result = SimpleNamespace( + sampled_tokens_cpu=sampled_tokens, + accepted_tokens_cpu=None, + active_request_ids=context.request_ids.long(), + finished_request_ids=torch.tensor([11]), + survivor_idxs=None, + newly_paused_request_ids=torch.tensor([10]), + evict_request_ids=torch.tensor([11]), + ) + input_ids = torch.empty(1, dtype=torch.int64) + position_ids = torch.empty(1, dtype=torch.int64) + call_order = [] + + controller._synchronize_async_sched_event = mock.Mock( + side_effect=lambda event: call_order.append(f"wait:{event}") + ) + controller._run_async_sched_sample = mock.Mock( + side_effect=lambda: call_order.append("sample") or sample_result + ) + + def update_request_state(*_args): + call_order.append("update") + context.num_prefill_requests = launched_prefill_requests + return request_result + + controller._run_async_sched_update_requests = mock.Mock(side_effect=update_request_state) + log_probs_gpu_result = object() + log_probs_transfer = SimpleNamespace(cpu_ready_event="log_probs") + controller._run_async_sched_log_probs = mock.Mock( + side_effect=lambda _: call_order.append("log_probs") or log_probs_gpu_result + ) + controller._copy_async_sched_log_probs_to_cpu = mock.Mock( + side_effect=lambda _: call_order.append("copy_log_probs") or log_probs_transfer + ) + controller._materialize_async_sched_log_probs = mock.Mock( + side_effect=lambda *_: call_order.append("materialize_log_probs") or ([[0.1], [0.2]], None) + ) + controller._dynamic_step_context_init = mock.Mock( + side_effect=lambda: call_order.append("context_init") or (input_ids, position_ids, None) + ) + controller._run_async_sched_forward = mock.Mock( + side_effect=lambda *_: call_order.append("forward") + ) + + async def yield_to_event_loop(_delay): + call_order.append("yield") + + with mock.patch( + "megatron.core.inference.text_generation_controllers." + "text_generation_controller.asyncio.sleep", + side_effect=yield_to_event_loop, + ): + result = asyncio.run( + controller._run_async_sched_step_no_overlap( + schedule_waiting_requests=lambda: call_order.append("admit") + ) + ) + + assert result.output["sample"].tolist() == [1, 2] + assert result.output["newly_paused_request_ids"].tolist() == [10] + assert result.output["evict_request_ids"].tolist() == [11] + assert result.decode_only == DecodeOnly( + consumed=consumed_prefill_requests == 0, launched=launched_prefill_requests == 0 + ) + assert result.output["log_probs"] == [[0.1], [0.2]] + assert call_order == [ + "sample", + "log_probs", + "copy_log_probs", + "wait:sample", + "update", + "admit", + "context_init", + "forward", + "wait:log_probs", + "materialize_log_probs", + "yield", + ] + context.resolve_requests.assert_not_called() + context.prepare_requests.assert_not_called() + context.commit_sampled_tokens.assert_not_called() + + +def test_async_sched_no_overlap_finishes_with_matching_ep_base_forward(): + """A rank that resolves its last request still matches the peer base collective.""" + context = _make_async_sched_context(total_request_count=1) + context.num_prefill_requests = 1 + model_config = SimpleNamespace( + params_dtype=torch.float32, + expert_model_parallel_size=2, + num_moe_experts=4, + moe_enable_routing_replay=False, + ) + controller = _make_async_sched_controller(context, model_config) + controller._async_sched_logits = AsyncScheduleLogitsState(is_valid=True) + sampled_tokens = torch.tensor([1]) + sample_result = SimpleNamespace( + sampled_tokens_gpu=sampled_tokens, + sampled_tokens_cpu_view=sampled_tokens, + sampled_mtp_tokens_gpu=None, + sampled_mtp_tokens_cpu_view=None, + accepted_tokens_cpu_view=None, + sample_cpu_ready_event=None, + ) + request_result = SimpleNamespace( + sampled_tokens_cpu=sampled_tokens, + accepted_tokens_cpu=None, + active_request_ids=context.request_ids.long(), + finished_request_ids=context.request_ids.clone(), + survivor_idxs=None, + newly_paused_request_ids=None, + evict_request_ids=None, + ) + controller._run_async_sched_sample = mock.Mock(return_value=sample_result) + controller._synchronize_async_sched_event = mock.Mock() + + def update_last_request(*_args): + context.total_request_count = 0 + context.active_token_count = 0 + return request_result + + controller._run_async_sched_update_requests = mock.Mock(side_effect=update_last_request) + controller._run_dummy_async_sched_base_step = mock.Mock() + controller._run_async_sched_forward = mock.Mock() + + result = asyncio.run( + controller._run_async_sched_step_no_overlap(schedule_waiting_requests=None) + ) + + assert result.output["finished_request_ids"].tolist() == [10] + controller._run_dummy_async_sched_base_step.assert_called_once_with() + controller._run_async_sched_forward.assert_not_called() + + +def test_async_sched_mtp_overlap_step_order(): + """MTP verification and rewind precede prepare while forward precedes resolve.""" + context = _make_async_sched_context(total_request_count=3) + controller = _make_async_sched_controller(context) + controller.num_speculative_tokens = 2 + controller._async_sched_logits = AsyncScheduleLogitsState( + is_valid=True, cuda_graph_request_count=7 + ) + sampled_tokens = torch.tensor([1, 4, 7]) + sampled_mtp_tokens = torch.tensor([[2, 5, 8], [3, 6, 9]]) + accepted_tokens = torch.tensor([9, 10, 11]) + sample_result = SimpleNamespace( + sampled_tokens_gpu=sampled_tokens, + sampled_tokens_cpu_view=sampled_tokens, + sampled_mtp_tokens_gpu=sampled_mtp_tokens, + sampled_mtp_tokens_cpu_view=sampled_mtp_tokens, + accepted_tokens_cpu_view=accepted_tokens, + accepted_counts_cpu_view=torch.tensor([1, 1, 1]), + sample_cpu_ready_event="sample", + ) + resolve_result = SimpleNamespace( + sampled_tokens_cpu=sampled_tokens, + accepted_tokens_cpu=accepted_tokens, + active_request_ids=context.request_ids.long(), + finished_request_ids=torch.empty(0, dtype=torch.int32), + survivor_idxs=torch.tensor([2, 1]), + newly_paused_request_ids=None, + evict_request_ids=None, + ) + input_ids = torch.empty(9, dtype=torch.int64) + position_ids = torch.empty(9, dtype=torch.int64) + call_order = [] + + controller._run_async_sched_sample_mtp = mock.Mock( + side_effect=lambda: call_order.append("sample_mtp") or sample_result + ) + controller._run_async_sched_mtp_rewind = mock.Mock( + side_effect=lambda *_args, **_kwargs: call_order.append("rewind") + ) + log_probs_gpu_result = object() + log_probs_transfer = SimpleNamespace(cpu_ready_event="log_probs") + controller._run_async_sched_log_probs = mock.Mock( + side_effect=lambda _: call_order.append("log_probs") or log_probs_gpu_result + ) + controller._copy_async_sched_log_probs_to_cpu = mock.Mock( + side_effect=lambda _: call_order.append("copy_log_probs") or log_probs_transfer + ) + controller._materialize_async_sched_log_probs = mock.Mock( + side_effect=lambda *_: call_order.append("materialize_log_probs") + or ([[0.1], [0.2], [0.3]], None) + ) + controller._run_async_sched_prepare = mock.Mock( + side_effect=lambda: call_order.append("prepare") or (input_ids, position_ids) + ) + context.copy_async_sched_sample_to_forward = mock.Mock( + side_effect=lambda *_: call_order.append("copy_input") + ) + controller._run_async_sched_publish_bookkeeping = mock.Mock( + side_effect=lambda: call_order.append("publish") or "bookkeeping" + ) + controller._run_async_sched_forward = mock.Mock( + side_effect=lambda *_: call_order.append("forward") + ) + controller._synchronize_async_sched_event = mock.Mock( + side_effect=lambda event: call_order.append(f"wait:{event}") + ) + context.commit_sampled_tokens = mock.Mock(side_effect=lambda *_: call_order.append("commit")) + controller._run_async_sched_resolve = mock.Mock( + side_effect=lambda *_: call_order.append("resolve") or resolve_result + ) + + result = asyncio.run(controller._run_async_sched_step_overlap_mtp()) + + assert result.output["accepted_tokens"].tolist() == [9, 10, 11] + assert result.decode_only == DecodeOnly(consumed=True, launched=True) + assert result.output["log_probs"] == [[0.1], [0.2], [0.3]] + assert call_order == [ + "sample_mtp", + "rewind", + "log_probs", + "copy_log_probs", + "prepare", + "copy_input", + "publish", + "forward", + "wait:sample", + "wait:bookkeeping", + "resolve", + "commit", + "wait:log_probs", + "materialize_log_probs", + ] + committed_tokens, committed_mtp_tokens = context.commit_sampled_tokens.call_args.args + assert torch.equal(committed_tokens, torch.tensor([7, 4])) + assert torch.equal(committed_mtp_tokens, torch.tensor([[8, 5], [9, 6]])) + + +@pytest.mark.parametrize( + "mode, run_async_overlap, has_pending_logits, num_speculative_tokens, " + "expected_method, expected_output", [ - (AsyncScheduleMode.LEGACY, 0, False, "legacy"), - (AsyncScheduleMode.SERIAL, 1, False, "legacy"), - (AsyncScheduleMode.SERIAL, 0, False, "async"), + (AsyncScheduleMode.LEGACY, None, True, 0, "legacy", "legacy"), + (AsyncScheduleMode.ASYNC, False, True, 0, "no_overlap", "no_overlap"), + (AsyncScheduleMode.ASYNC, True, False, 0, "no_overlap", "no_overlap"), + (AsyncScheduleMode.ASYNC, True, True, 0, "overlap", "overlap"), + (AsyncScheduleMode.ASYNC, True, True, 2, "overlap_mtp", "overlap_mtp"), ], ) def test_async_generate_output_tokens_dynamic_batch_routes( - mode, num_prefill_requests, skip_bookkeeping, expected_result + mode, + run_async_overlap, + has_pending_logits, + num_speculative_tokens, + expected_method, + expected_output, ): context = _make_async_sched_context() context.config.async_sched_mode = mode - context.num_prefill_requests = num_prefill_requests controller = _make_async_sched_controller(context) - controller._run_legacy_step = mock.AsyncMock(return_value="legacy") - controller._run_async_sched_serial_step = mock.AsyncMock(return_value="async") - - result = asyncio.run(controller.async_generate_output_tokens_dynamic_batch(skip_bookkeeping)) + controller._async_sched_logits.is_valid = has_pending_logits + controller.num_speculative_tokens = num_speculative_tokens + controller._validate_async_sched_support_for_step = mock.Mock() + controller._run_legacy_step = mock.AsyncMock( + return_value=DynamicBatchControllerStepResult( + decode_only=DecodeOnly(consumed=True, launched=True), output="legacy" + ) + ) + controller._run_async_sched_step_no_overlap = mock.AsyncMock( + return_value=DynamicBatchControllerStepResult( + decode_only=DecodeOnly(consumed=False, launched=True), output="no_overlap" + ) + ) + controller._run_async_sched_step_overlap = mock.AsyncMock( + return_value=DynamicBatchControllerStepResult( + decode_only=DecodeOnly(consumed=True, launched=True), output="overlap" + ) + ) + controller._run_async_sched_step_overlap_mtp = mock.AsyncMock( + return_value=DynamicBatchControllerStepResult( + decode_only=DecodeOnly(consumed=True, launched=True), output="overlap_mtp" + ) + ) + schedule_waiting_requests = mock.Mock() + + kwargs = ( + {} + if run_async_overlap is None + else { + "run_async_overlap": run_async_overlap, + "schedule_waiting_requests": schedule_waiting_requests, + } + ) + result = asyncio.run(controller.async_generate_output_tokens_dynamic_batch(**kwargs)) - assert result == expected_result + assert result.output == expected_output + methods = { + "legacy": controller._run_legacy_step, + "no_overlap": controller._run_async_sched_step_no_overlap, + "overlap": controller._run_async_sched_step_overlap, + "overlap_mtp": controller._run_async_sched_step_overlap_mtp, + } + assert methods[expected_method].await_count == 1 @pytest.mark.parametrize( "mode, expected_message", [ - (AsyncScheduleMode.SERIAL, "request bookkeeping"), + (AsyncScheduleMode.ASYNC, "request bookkeeping"), ("unexpected", "Unexpected async scheduling mode"), ], ) @@ -512,7 +1636,9 @@ def test_async_generate_output_tokens_dynamic_batch_assertions(mode, expected_me context.config.async_sched_mode = mode controller = _make_async_sched_controller(context) controller._run_legacy_step = mock.AsyncMock() - controller._run_async_sched_serial_step = mock.AsyncMock() + controller._run_async_sched_step_no_overlap = mock.AsyncMock() + controller._run_async_sched_step_overlap = mock.AsyncMock() + controller._run_async_sched_step_overlap_mtp = mock.AsyncMock() with pytest.raises(AssertionError, match=expected_message): asyncio.run(controller.async_generate_output_tokens_dynamic_batch(skip_bookkeeping=True)) @@ -533,6 +1659,84 @@ def teardown_class(cls): def teardown_method(self, method): InferenceMode.unset_active() + @pytest.mark.internal + def test_async_sched_no_overlap_pauses_boundary_request(self): + """No-overlap uses real lifecycle bookkeeping before forwarding survivors.""" + self.setup_model( + torch.float32, batch_size=2, static=False, block_size_tokens=4, max_requests=2 + ) + controller = self.text_generation_controller + context = controller.inference_wrapped_model.inference_context + context.reset() + + active_slice = slice(0, 2) + context.total_request_count = 2 + context.active_token_count = 2 + context.request_ids[active_slice] = torch.tensor([10, 11], dtype=torch.int32) + context.request_in_prefill_status_tensor[active_slice] = 0 + context.request_query_lengths[active_slice] = 1 + context.request_output_lengths[active_slice] = 16 + context.request_kv_length_offsets[active_slice] = 3 + context.request_last_kv_block_offset[active_slice] = torch.tensor( + [context.block_size_tokens - 1, 0], dtype=torch.int32 + ) + context.request_metadata["termination_id"][active_slice] = 99 + context.build_active_slices(2) + + block_ids = context.kv_block_allocator.allocate_memory_blocks(2) + context.request_to_kv_block_ids[active_slice, 0] = block_ids + context.request_last_kv_block_id[active_slice] = block_ids + context.request_kv_block_counts[active_slice] = 1 + context.token_to_input_ids[active_slice] = torch.tensor([80, 81]) + + # Leave room in paused storage but no capacity to keep both requests active. + context.kv_block_allocator.active_count = context.kv_block_allocator.get_active_used() + context.kv_block_allocator.paused_count = 2 + context.kv_block_allocator.total_avail = 0 + + sampled_tokens = torch.tensor([90, 91], dtype=torch.int64) + controller._async_sched_logits = AsyncScheduleLogitsState( + is_valid=True, cuda_graph_request_count=2 + ) + controller._run_async_sched_sample = mock.Mock( + return_value=SimpleNamespace( + sampled_tokens_gpu=sampled_tokens, + sampled_tokens_cpu_view=sampled_tokens, + sampled_mtp_tokens_gpu=None, + sampled_mtp_tokens_cpu_view=None, + accepted_tokens_cpu_view=None, + sample_cpu_ready_event=None, + ) + ) + forward_input_ids = torch.tensor([91]) + forward_position_ids = torch.tensor([4]) + + def initialize_survivor_forward(): + assert context.paused_request_count == 1 + assert context.request_ids[:2].tolist() == [10, 11] + assert context.token_to_input_ids[0].item() == 91 + return forward_input_ids, forward_position_ids, None + + controller._dynamic_step_context_init = mock.Mock(side_effect=initialize_survivor_forward) + controller._run_async_sched_forward = mock.Mock() + + result = asyncio.run( + controller._run_async_sched_step_no_overlap(schedule_waiting_requests=None) + ).output + + assert result["sample"].tolist() == [90, 91] + assert result["finished_request_ids"].numel() == 0 + assert result["newly_paused_request_ids"].flatten().tolist() == [10] + assert result["evict_request_ids"] is None + assert context.paused_request_count == 1 + active_request_ids = context.request_ids[ + context.paused_request_count : context.total_request_count + ] + assert active_request_ids.tolist() == [11] + controller._run_async_sched_forward.assert_called_once_with( + forward_input_ids, forward_position_ids + ) + def test_sample_from_logits(self): self.setup_model(torch.float32) @@ -628,7 +1832,7 @@ def test_sample_from_dynamic_logits( ): if backend == "flashinfer": pytest.importorskip("flashinfer") - batch_size = 15 + batch_size = 18 self.setup_model( torch.float32, batch_size=batch_size, @@ -649,6 +1853,7 @@ def test_sample_from_dynamic_logits( (SamplingParams(top_p=0.8), [4, 1, 7]), (SamplingParams(temperature=10.0, top_k=5), [11, 5, 8]), (SamplingParams(temperature=0.0, top_k=1), [12, 13, 14]), + (SamplingParams(temperature=1.2), [15, 16, 17]), ] # For non-torch backends, test simultaneous top_k and top_p sampling. if backend != "torch": @@ -711,6 +1916,71 @@ def test_sample_from_dynamic_logits( sampled_logits >= expected_min_values ), f"The sampled logits should all be greater than {expected_min_values} but its {sampled_logits}" + def test_dynamic_sampling_keeps_sampled_tokens_buffer_full_capacity(self): + """`_sampled_tokens_cuda` is a single `max_requests` buffer written in place by + every sampling path. The non-speculative path (`_dynamic_step_sample_logits`) + writes its `active_request_count` prefix, and the async-scheduling path + (`_run_async_sched_sample`) writes its own prefix through `sample_kernel`. The buffer + must retain its full capacity across successive steps regardless of each step's + active count, so a later step with more active requests than an earlier one still + has an in-bounds destination. + + Drive a small-batch non-speculative sample followed by a larger-batch async + sample through the same buffer, and confirm the buffer keeps its capacity and + identity and the larger async write lands correctly. + """ + self.setup_model( + torch.float32, batch_size=8, static=False, materialize_only_last_token_logits=True + ) + context = self.text_generation_controller.inference_wrapped_model.inference_context + controller = self.text_generation_controller + + capacity = controller._sampled_tokens_cuda.numel() + buffer_ptr = controller._sampled_tokens_cuda.data_ptr() + small_count, large_count = 2, 5 + assert large_count < capacity + + # Non-speculative sample over a small active batch. top_k == 1 makes the torch + # backend short-circuit to argmax, so the sampled token is deterministic. + context.active_request_metadata["temperature"][:small_count].fill_(1.0) + context.active_request_metadata["top_k"][:small_count].fill_(1) + context.active_request_metadata["top_p"][:small_count].fill_(0.0) + context.padded_active_token_count = small_count + context.request_query_lengths = torch.ones(small_count, dtype=torch.int32, device="cuda") + context.paused_request_count = 0 + context.total_request_count = small_count + context.num_prefill_requests = 0 + context.pad_active_slices() + + small_expected = torch.tensor([3, 4], device="cuda") + small_logits = torch.zeros(1, small_count, self.vocab_size, device="cuda") + for row, col in enumerate(small_expected.tolist()): + small_logits[0, row, col] = 10.0 + controller._all_logits_cuda = small_logits + controller._dynamic_step_sample_logits() + + assert controller._sampled_tokens_cuda.numel() == capacity + assert controller._sampled_tokens_cuda.data_ptr() == buffer_ptr + assert torch.equal(controller._sampled_tokens_cuda[:small_count], small_expected) + + # Async-scheduling sample over a larger active batch through the same buffer. + context.total_request_count = large_count + context.paused_request_count = 0 + context.active_request_metadata["temperature"][:large_count].fill_(1.0) + context.active_request_metadata["top_k"][:large_count].fill_(1) + context.active_request_metadata["top_p"][:large_count].fill_(0.0) + large_expected = torch.tensor([0, 1, 2, 3, 4], device="cuda") + large_logits = torch.zeros(1, large_count, self.vocab_size, device="cuda") + for row, col in enumerate(large_expected.tolist()): + large_logits[0, row, col] = 10.0 + controller._all_logits_cuda = large_logits + + sampled_tokens_gpu = controller._run_async_sched_sample().sampled_tokens_gpu + + assert sampled_tokens_gpu.data_ptr() == controller._sampled_tokens_cuda.data_ptr() + assert controller._sampled_tokens_cuda.numel() == capacity + assert torch.equal(sampled_tokens_gpu, large_expected) + @pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) @pytest.mark.parametrize( "symmetric_ar_type", @@ -793,7 +2063,7 @@ def test_generate_all_output_tokens_static_batch(self, dtype, symmetric_ar_type, assert ( len(request.segments) == len(request.prompt_log_probs) + len(request.generated_log_probs) + 1 - ), "Segments should be returned for both prompt and generated tokens" + ), f"Segments should be returned for both prompt and generated tokens: {request}" assert len(request.prompt) + len(request.generated_text) == len( request.text ), "Output text should include prompts and generations" @@ -1861,7 +3131,7 @@ def test_mtp_sp_padding_real_ranks(self, active_request_count): assert not hasattr(unwrapped_model, '_decoder_hidden_states_cache') def test_mtp_sp_padding_dummy_ranks(self): - """Test _dummy_serial_mtp_forward with real MTP layers and sequence parallelism. + """Test _run_dummy_serial_mtp_forward with real MTP layers and sequence parallelism. Creates a GPTModel with real MTP layers and SP, then runs the dummy forward path used by EP dummy ranks. Verifies the full MTP forward @@ -1887,10 +3157,10 @@ def test_mtp_sp_padding_dummy_ranks(self): unwrapped_model._decoder_hidden_states_cache = True # Run the dummy MTP forward path end-to-end. - ctrl._dummy_serial_mtp_forward() + ctrl._run_dummy_serial_mtp_forward() # Verify compute_mtp_single_step produces correctly-shaped outputs - # with the same dummy tensor shapes that _dummy_serial_mtp_forward uses. + # with the same dummy tensor shapes that _run_dummy_serial_mtp_forward uses. # padded_count == tp_size when SP is enabled. dummy_hidden = torch.zeros((1, 1, self.hidden_size), device='cuda', dtype=torch.float32) dummy_tokens = torch.zeros((1, tp_size), device='cuda', dtype=torch.long) @@ -1990,6 +3260,7 @@ def setup_model( Utils.initialize_model_parallel( tensor_model_parallel_size=tensor_model_parallel_size, pipeline_model_parallel_size=pipeline_model_parallel_size, + expert_model_parallel_size=expert_model_parallel_size, ) super().setup_model( dtype, @@ -2159,3 +3430,121 @@ def test_sampled_tokens_match_with_parallelism(self, static, tp_size, pp_size): assert ( expected == actual ), f"Rank {i} tokens differ from rank {local_rank} tokens for request {j}" + + @pytest.mark.parametrize("static", [True, False]) + @pytest.mark.parametrize("enable_prefix_caching", [True, False]) + def test_sampled_tokens_dp_mismatch(self, static, enable_prefix_caching): + """ + TextGenerationController should generate different tokens + on every DP rank given the same prompt / request. + """ + if not static and not is_fa_min_version("2.7.3"): + pytest.skip(reason="Need latest flash attn for dynamic batching") + + self.setup_model( + dtype=torch.bfloat16, + # Set all model parallelisms to 1. + # No rank should shard the model. + tensor_model_parallel_size=1, + pipeline_model_parallel_size=1, + expert_model_parallel_size=1, + # Test all batching and balancing strategies. + static=static, + enable_prefix_caching=enable_prefix_caching, + # Set a random seed for generation, so we can + # verify that different DP ranks produce disparate + # generations given the DP rank seed offset. + # Without this, generations will always be random + # and we don't be able to verify DP disparity. + use_training_random_init=True, + ) + # Ensure only data parallelism is used. This is critical since + # we expect generation parity across model parallel ranks. + dp_size = parallel_state.get_data_parallel_group().size() + assert ( + dp_size == torch.distributed.get_world_size() + ), "[test_sampled_tokens_dp_mismatch] Expected DP size to match WORLD size for this DP disparity test." + + # Prepare requests. + active_requests: Dict[str, InferenceRequest] = OrderedDict() + for i in range(self.batch_size): + # Create a batch of constant prompts to test DP disparity. + # Same inputs should produce different outputs. + prompt = "sample" * (i + 1) + prompt_tokens = [1] * (i + 1) + request_id = str(i) + inference_request = InferenceRequest( + request_id=request_id, + prompt=prompt, + sampling_params=SamplingParams( + top_k=10, num_tokens_to_generate=25, return_log_probs=True + ), + arrival_time=time.time(), + prompt_tokens=prompt_tokens, + status=Status.ACTIVE_BUT_NOT_GENERATING_TOKENS, + ) + active_requests[request_id] = inference_request + + # Generate tokens for each sample of the batch. + if static: + # Static batching requires dummy metadata and functions. + self.mock_tokenizer.vocab_size = self.vocab_size + self.mock_tokenizer.eod = self.vocab_size - 1 + self.mock_tokenizer.detokenize.side_effect = ( + lambda x, skip_special_tokens=False: ' '.join( + [ + ''.join(random.choices(string.ascii_letters, k=random.randint(4, 10))) + for _ in range(len(x)) + ] + ) + ) + self.mock_tokenizer.offsets.side_effect = lambda _, s: [ + i for i, c in enumerate(s) if c == ' ' + ] + [len(s)] + + # Generate. + requests = self.text_generation_controller.generate_all_output_tokens_static_batch( + active_requests + ) + all_generated_tokens = [req.generated_tokens.tolist() for req in requests.values()] + else: + all_generated_tokens = [[] for _ in range(len(active_requests))] + context = self.text_generation_controller.inference_wrapped_model.inference_context + for request_id, request in active_requests.items(): + context.add_request( + DynamicInferenceRequest( + request_id=int(request_id), + prompt_tokens=torch.tensor( + request.prompt_tokens, + dtype=torch.long, + device=torch.cuda.current_device(), + ), + sampling_params=SamplingParams( + top_k=10, return_log_probs=True, num_tokens_to_generate=25 + ), + ) + ) + expected_active_requests = set(int(x) for x in active_requests.keys()) + while context.has_unfinished_requests(): + result = self.text_generation_controller.generate_output_tokens_dynamic_batch() + new_tokens = result["sample"] + active_ids = result["active_request_ids"].tolist() + finished_ids = result["finished_request_ids"].tolist() + assert len(new_tokens) == len(expected_active_requests) + assert set(active_ids) == expected_active_requests + expected_active_requests -= set(finished_ids) + for i, token in enumerate(new_tokens.tolist()): + all_generated_tokens[i].append(token) + + # Wait for all requests on all host ranks to complete before proceeding. + torch.distributed.barrier() + + # All-gather the generated tokens on every DP rank. + all_dp_generated_tokens = [None] * dp_size + torch.distributed.all_gather_object(all_dp_generated_tokens, all_generated_tokens) + for i in range(self.batch_size): + # Get the i-th generation for each DP rank. + dp_batch = [tuple(batch[i]) for batch in all_dp_generated_tokens] + assert len(set(dp_batch)) == len( + dp_batch + ), "Detected duplicate generations across DP ranks." diff --git a/tests/unit_tests/models/mimo/test_mimo_1f1b_schedule.py b/tests/unit_tests/models/mimo/test_mimo_1f1b_schedule.py index bbc665bc44d..56d8db19735 100644 --- a/tests/unit_tests/models/mimo/test_mimo_1f1b_schedule.py +++ b/tests/unit_tests/models/mimo/test_mimo_1f1b_schedule.py @@ -7,7 +7,6 @@ """ import logging -from contextlib import ExitStack, contextmanager from functools import partial from types import SimpleNamespace @@ -61,27 +60,6 @@ _embedding_pg_cache: dict = {} -def build_no_sync_func(mimo_model): - """Build a no_sync_func that stacks DDP no_sync over each sub-module. - - Shared by 1F1B pipeline tests and colocated-correctness tests — both need - DDP's gradient sync disabled during microbatches and resumed via the - schedule's finalize_grads_func. - """ - - @contextmanager - def no_sync_func(): - with ExitStack() as stack: - if mimo_model.language_model is not None: - stack.enter_context(mimo_model.language_model.no_sync()) - for submodule in mimo_model.modality_submodules.values(): - if submodule is not None: - stack.enter_context(submodule.no_sync()) - yield - - return no_sync_func - - def create_hypercomm_grid(offset=0, tp=1, cp=1, pp=1, dp=1): """Create a HyperCommGrid (base view) plus a dense expert view, matching the topology builder. @@ -623,8 +601,6 @@ def run_mimo_1f1b_test( per_token_loss=True, ) - mimo_model.config.no_sync_func = build_no_sync_func(mimo_model) - # Use the production grad-sync hook (finalize per module over its own groups + # cross-grid N_global per-token mean) for every config. grad_sync_topology = SimpleNamespace( @@ -634,7 +610,11 @@ def run_mimo_1f1b_test( **{name: vision_pg for name in mimo_model.modality_submodules}, }, ) - configure_grad_sync(SimpleNamespace(), mimo_model, grad_sync_topology) + configure_grad_sync( + SimpleNamespace(overlap_grad_reduce=True, align_grad_reduce=False), + mimo_model, + grad_sync_topology, + ) # Create optimizer opt_config = OptimizerConfig( diff --git a/tests/unit_tests/models/mimo/test_mimo_colocated_correctness.py b/tests/unit_tests/models/mimo/test_mimo_colocated_correctness.py index 747b66a815a..251a4228740 100644 --- a/tests/unit_tests/models/mimo/test_mimo_colocated_correctness.py +++ b/tests/unit_tests/models/mimo/test_mimo_colocated_correctness.py @@ -68,7 +68,6 @@ from megatron.core.transformer.enums import ModelType from megatron.core.utils import unwrap_model from tests.unit_tests.models.mimo.test_mimo_1f1b_schedule import ( - build_no_sync_func, create_all_embedding_groups, create_hypercomm_grid, destroy_all_grids, @@ -167,7 +166,7 @@ def _set_deterministic_env(): def _wire_training_hooks(mimo_model, module_to_grid_map, language_pg, vision_pg): - """Attach no_sync plus the production grad-sync hooks to a MimoModel. + """Attach the production no-sync and grad-sync hooks to a MimoModel. Delegates the finalize/grad-scale wiring to ``configure_grad_sync`` (the real examples/mimo path), so this test's dp1-reference assertions validate that @@ -175,7 +174,6 @@ def _wire_training_hooks(mimo_model, module_to_grid_map, language_pg, vision_pg) mean: all-reduce ``total_num_tokens`` over the LLM DP group to get ``N_global``, finalize each submodule over its own group, then ``scale_gradients(1/N_global)``. """ - mimo_model.config.no_sync_func = build_no_sync_func(mimo_model) topology = SimpleNamespace( grids=module_to_grid_map, module_pgs={ @@ -183,7 +181,9 @@ def _wire_training_hooks(mimo_model, module_to_grid_map, language_pg, vision_pg) **{name: vision_pg for name in mimo_model.modality_submodules}, }, ) - configure_grad_sync(SimpleNamespace(), mimo_model, topology) + configure_grad_sync( + SimpleNamespace(overlap_grad_reduce=True, align_grad_reduce=False), mimo_model, topology + ) def _generate_and_broadcast_global_batches( diff --git a/tests/unit_tests/models/mimo/test_mimo_encoder_prefetch.py b/tests/unit_tests/models/mimo/test_mimo_encoder_prefetch.py new file mode 100644 index 00000000000..a53c682e2fb --- /dev/null +++ b/tests/unit_tests/models/mimo/test_mimo_encoder_prefetch.py @@ -0,0 +1,500 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +from __future__ import annotations + +import threading +import time +from contextlib import nullcontext +from types import SimpleNamespace + +import pytest +import torch +from torch import nn + +from examples.mimo.training import encoder_prefetch +from examples.mimo.training import step as mimo_step +from examples.mimo.training.encoder_prefetch import ( + PREFETCHED_FEATURES_KEY, + PROJECTION_TIMER_KEY, + EncoderPrefetchLoader, + prefetch_frozen_features, + validate_encoder_prefetch_args, +) +from megatron.core.models.mimo.submodules.vision import VisionModalitySubmodules + +ENCODER = "clip_encoder" + + +def _args(**overrides): + values = { + "mimo_encoder_prefetch": True, + "mimo_encoder_prefetch_depth": 2, + "freeze_vit": True, + "freeze_projection": False, + "encoder_tp": 1, + "rerun_mode": "disabled", + } + values.update(overrides) + return SimpleNamespace(**values) + + +@pytest.mark.parametrize( + ("field", "value", "match"), + [ + ("freeze_vit", False, "freeze-vit"), + ("freeze_projection", True, "trainable projection"), + ("encoder_tp", 2, "TP=1"), + ("encoder_cp", 2, "CP=1"), + ("encoder_pp", 2, "PP=1"), + ("encoder_ep", 2, "EP=1"), + ("mimo_encoder_prefetch_depth", 0, "positive"), + ("rerun_mode", "validate_results", "rerun"), + ], +) +def test_prefetch_validation(field, value, match): + with pytest.raises(ValueError, match=match): + validate_encoder_prefetch_args(_args(**{field: value})) + + +class _LinearEncoder(nn.Module): + def __init__(self): + super().__init__() + self.linear = nn.Linear(4, 4, bias=False) + + def forward(self, *, x): + return self.linear(x) + + +def test_prefetch_skips_backbone_autograd_but_keeps_projection_gradients(): + torch.manual_seed(123) + submodule = VisionModalitySubmodules( + encoders={"radio": _LinearEncoder()}, input_projections=[nn.Linear(4, 4, bias=False)] + ) + submodule.encoders.requires_grad_(False) + inputs = {"radio": {"x": torch.randn(2, 3, 4)}} + + expected = submodule(encoder_inputs=inputs) + features = prefetch_frozen_features(submodule, inputs) + actual = submodule(hidden_states=features) + + torch.testing.assert_close(actual, expected) + actual.sum().backward() + assert all(parameter.grad is None for parameter in submodule.encoders.parameters()) + assert all(parameter.grad is not None for parameter in submodule.input_projections.parameters()) + with pytest.raises(ValueError, match="mutually exclusive"): + submodule(encoder_inputs=inputs, hidden_states=features) + + +class _FakeEvent: + def __init__(self): + self.recorded_on = None + self.record_calls = 0 + self.synchronized = False + + def record(self, stream=None): + self.recorded_on = stream + self.record_calls += 1 + + def synchronize(self): + self.synchronized = True + + def query(self): + return self.record_calls > 0 + + def elapsed_time(self, _end_event): + return 1.25 + + +class _FakeStream: + def __init__(self): + self.waited_events = [] + self.synchronize_calls = 0 + + def wait_event(self, event): + self.waited_events.append(event) + + def synchronize(self): + self.synchronize_calls += 1 + + +@pytest.fixture +def fake_cuda(monkeypatch): + current = _FakeStream() + producer = _FakeStream() + events = [] + + def make_event(**_kwargs): + event = _FakeEvent() + events.append(event) + return event + + monkeypatch.setattr(torch.cuda, "current_device", lambda: 0) + monkeypatch.setattr(torch.cuda, "device", lambda _device: nullcontext()) + monkeypatch.setattr(torch.cuda, "current_stream", lambda: current) + monkeypatch.setattr(torch.cuda, "Event", make_event) + monkeypatch.setattr(torch.cuda, "stream", lambda _stream: nullcontext()) + monkeypatch.setattr(encoder_prefetch, "move_batch_to_cuda", lambda value: value) + return SimpleNamespace(current=current, producer=producer, events=events) + + +class _Source: + def __init__(self, count): + self.count = count + self.position = 0 + self.condition = threading.Condition() + + def __iter__(self): + return self + + def __next__(self): + with self.condition: + if self.position == self.count: + raise StopIteration + sequence = self.position + self.position += 1 + self.condition.notify_all() + return { + "input_ids": torch.tensor(sequence), + "modality_inputs": {ENCODER: {"radio": {"x": torch.tensor(sequence)}}}, + } + + +def _wait_until(predicate): + deadline = time.monotonic() + 2 + while not predicate(): + assert time.monotonic() < deadline + time.sleep(0.005) + + +@pytest.mark.parametrize("depth", (1, 2, 4)) +def test_depth_is_completed_batches_and_refills_after_pop(fake_cuda, depth): + source = _Source(8) + produced = [] + + def produce(inputs): + produced.append(int(inputs["radio"]["x"].item())) + return torch.tensor(produced[-1]) + + loader = EncoderPrefetchLoader( + source=source, + encoder_name=ENCODER, + feature_producer=produce, + depth=depth, + stream=fake_cuda.producer, + ) + loader.start() + _wait_until(lambda: len(loader._ready) == depth) + + first = next(loader) + _wait_until(lambda: len(loader._ready) == depth) + + assert first[PREFETCHED_FEATURES_KEY][ENCODER].item() == 0 + assert PROJECTION_TIMER_KEY not in first + assert len(fake_cuda.producer.waited_events) == 1 + assert fake_cuda.current.waited_events == [] + loader.close() + + +def test_pending_encode_can_be_claimed_while_cpu_read_ahead_blocks(fake_cuda, caplog): + caplog.set_level("INFO", logger=f"{encoder_prefetch.__name__}.debug") + read_ahead_started = threading.Event() + release_read_ahead = threading.Event() + + class _BlockingSource(_Source): + def __next__(self): + if self.position == 1: + read_ahead_started.set() + release_read_ahead.wait(timeout=2) + return super().__next__() + + loader = EncoderPrefetchLoader( + source=_BlockingSource(2), + encoder_name=ENCODER, + feature_producer=lambda inputs: torch.tensor(inputs["radio"]["x"].item()), + depth=1, + stream=fake_cuda.producer, + debug=True, + ) + loader.start() + assert read_ahead_started.wait(timeout=1) + + try: + assert len(loader._ready) == 0 + completion_event = loader._pending[1] + batch = next(loader) + assert batch[PREFETCHED_FEATURES_KEY][ENCODER].item() == 0 + assert fake_cuda.current.waited_events[-1] is completion_event + assert loader._pending is None + finally: + release_read_ahead.set() + loader.close() + + assert "consumer-wait batch=0 encoder_wait_ms=1.250" in caplog.text + assert "claimed_pending=1" in caplog.text + + +def test_cpu_read_ahead_overlaps_encode_without_enqueuing_another_batch(fake_cuda): + class _ObservedSource(_Source): + def __next__(self): + if self.position == 1: + assert not fake_cuda.events[-1].synchronized + return super().__next__() + + source = _ObservedSource(2) + produced = [] + + def produce(inputs): + produced.append(inputs["radio"]["x"].item()) + return torch.tensor(produced[-1]) + + loader = EncoderPrefetchLoader( + source=source, + encoder_name=ENCODER, + feature_producer=produce, + depth=1, + stream=fake_cuda.producer, + ) + loader.start() + _wait_until(lambda: len(loader._ready) == 1) + + assert source.position == 2 + assert produced == [0] + + for expected in (0, 1): + batch = next(loader) + assert batch[PREFETCHED_FEATURES_KEY][ENCODER].item() == expected + with pytest.raises(StopIteration): + next(loader) + loader.close() + + +def test_prefetch_keeps_input_ids_on_cpu_path(fake_cuda, monkeypatch): + input_ids = torch.tensor([[511, 1]]) + encoder_inputs = {"radio": {"x": torch.tensor(0)}} + moved = [] + + def record_move(value): + moved.append(value) + return value + + monkeypatch.setattr(encoder_prefetch, "move_batch_to_cuda", record_move) + loader = EncoderPrefetchLoader( + source=[{"input_ids": input_ids, "modality_inputs": {ENCODER: encoder_inputs}}], + encoder_name=ENCODER, + feature_producer=lambda _inputs: torch.tensor(0), + depth=1, + stream=fake_cuda.producer, + ) + loader.start() + _wait_until(lambda: len(loader._ready) == 1) + + batch = next(loader) + loader.close() + + assert batch["input_ids"] is input_ids + assert len(moved) == 1 + assert moved[0] is encoder_inputs + + +def test_debug_logs_prefetch_timing_and_queue_state(fake_cuda, caplog): + module_logger_level = encoder_prefetch.logger.level + caplog.set_level("INFO", logger=f"{encoder_prefetch.__name__}.debug") + loader = EncoderPrefetchLoader( + source=_Source(2), + encoder_name=ENCODER, + feature_producer=lambda inputs: torch.tensor(inputs["radio"]["x"].item()), + depth=1, + stream=fake_cuda.producer, + debug=True, + ) + loader.start() + _wait_until(lambda: len(loader._ready) == 1) + + first = next(loader) + with first.pop(PROJECTION_TIMER_KEY): + pass + _wait_until(lambda: len(loader._ready) == 1) + second = next(loader) + with second.pop(PROJECTION_TIMER_KEY): + pass + with pytest.raises(StopIteration): + next(loader) + loader.close() + + assert encoder_prefetch.logger.level == module_logger_level + assert "encoder-prefetch-debug consumer batch=0 ready_at_request=1/1" in caplog.text + assert "encoder-prefetch-debug producer batch=1" in caplog.text + assert "encoder-prefetch-debug projection batch=0 projection_ms=1.250" in caplog.text + + +def test_producer_failure_is_terminal_and_preserves_ready_fifo(fake_cuda): + class _FailingSource(_Source): + def __next__(self): + if self.position == 2: + raise ValueError("boom") + return super().__next__() + + loader = EncoderPrefetchLoader( + source=_FailingSource(8), + encoder_name=ENCODER, + feature_producer=lambda inputs: torch.tensor(inputs["radio"]["x"].item()), + depth=2, + stream=fake_cuda.producer, + ) + loader.start() + _wait_until(lambda: len(loader._ready) == 2) + + for expected in (0, 1): + batch = next(loader) + assert batch[PREFETCHED_FEATURES_KEY][ENCODER].item() == expected + with pytest.raises(RuntimeError, match="producer failed") as exc_info: + next(loader) + assert isinstance(exc_info.value.__cause__, ValueError) + loader.close() + + +def test_close_does_not_raise_when_worker_is_stuck(fake_cuda, caplog): + entered = threading.Event() + release = threading.Event() + + def produce(_inputs): + entered.set() + release.wait(timeout=2) + return torch.tensor(1) + + loader = EncoderPrefetchLoader( + source=_Source(1), + encoder_name=ENCODER, + feature_producer=produce, + depth=1, + stream=fake_cuda.producer, + worker_join_timeout_s=0.01, + ) + loader.start() + assert entered.wait(timeout=1) + + loader.close() + + assert "worker did not stop" in caplog.text + assert fake_cuda.producer.synchronize_calls == 0 + release.set() + assert loader._worker is not None + loader._worker.join(timeout=1) + assert not loader._worker.is_alive() + + +def test_forward_step_projects_prefetched_features_inside_debug_timer(monkeypatch): + events = [] + + class _Lease: + def __enter__(self): + events.append("enter") + + def __exit__(self, *_args): + events.append("exit") + + class _Model: + role = SimpleNamespace(modality_module_names=(ENCODER,)) + + def _forward_encoders(self, input_ids, modality_inputs, input_tensors): + assert modality_inputs is None + events.append("project") + return input_tensors + + features = {ENCODER: torch.ones(2, 4)} + batch = { + "input_ids": torch.tensor([[511]]), + PREFETCHED_FEATURES_KEY: features, + PROJECTION_TIMER_KEY: _Lease(), + } + monkeypatch.setattr( + mimo_step, + "move_batch_to_cuda", + lambda _value: pytest.fail("prefetched batches are already CUDA-resident"), + ) + + output, _ = mimo_step.mimo_forward_step(iter([batch]), _Model()) + + assert output is features + assert events == ["enter", "project", "exit"] + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_real_cuda_pending_handoff_waits_for_encode(): + read_ahead_started = threading.Event() + release_read_ahead = threading.Event() + backbone = nn.Linear(4, 4, bias=False, device="cuda") + + class _BlockingSource: + def __init__(self): + self.position = 0 + + def __iter__(self): + return self + + def __next__(self): + if self.position == 1: + read_ahead_started.set() + release_read_ahead.wait(timeout=2) + raise StopIteration + self.position += 1 + return { + "input_ids": torch.tensor([[511]]), + "modality_inputs": {ENCODER: {"radio": {"x": torch.ones(32, 4)}}}, + } + + def produce(inputs): + with torch.no_grad(): + return backbone(inputs["radio"]["x"]) + + loader = EncoderPrefetchLoader( + source=_BlockingSource(), encoder_name=ENCODER, feature_producer=produce, depth=1 + ) + loader.start() + assert read_ahead_started.wait(timeout=1) + + try: + assert len(loader._ready) == 0 + batch = next(loader) + expected = backbone(torch.ones(32, 4, device="cuda")) + torch.testing.assert_close(batch[PREFETCHED_FEATURES_KEY][ENCODER], expected) + finally: + release_read_ahead.set() + loader.close() + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_real_cuda_handoff_and_projection_gradients(): + backbone = nn.Linear(4, 4, bias=False, device="cuda") + projection = nn.Linear(4, 1, bias=False, device="cuda") + backbone.requires_grad_(False) + source = [ + { + "input_ids": torch.tensor([[sequence]]), + "modality_inputs": {ENCODER: {"radio": {"x": torch.full((32, 4), float(sequence))}}}, + } + for sequence in range(8) + ] + + def produce(inputs): + with torch.no_grad(): + return backbone(inputs["radio"]["x"]) + + loader = EncoderPrefetchLoader( + source=source, encoder_name=ENCODER, feature_producer=produce, depth=2 + ) + loader.start() + losses = [] + for sequence in range(8): + batch = next(loader) + assert not batch["input_ids"].is_cuda + features = batch[PREFETCHED_FEATURES_KEY][ENCODER] + reference = backbone(torch.full((32, 4), float(sequence), device="cuda")) + torch.testing.assert_close(features, reference) + losses.append(projection(features).sum()) + torch.stack(losses).sum().backward() + loader.close() + + assert projection.weight.grad is not None + assert torch.isfinite(projection.weight.grad).all() + assert all(parameter.grad is None for parameter in backbone.parameters()) diff --git a/tests/unit_tests/models/mimo/test_mimo_forward_step.py b/tests/unit_tests/models/mimo/test_mimo_forward_step.py index d6f470f8a82..21bbd94be0a 100644 --- a/tests/unit_tests/models/mimo/test_mimo_forward_step.py +++ b/tests/unit_tests/models/mimo/test_mimo_forward_step.py @@ -7,7 +7,8 @@ import pytest import torch -from examples.mimo.training.step import loss_func, move_batch_to_cuda +from examples.mimo.training.batch import move_batch_to_cuda +from examples.mimo.training.step import loss_func from megatron.core.packed_seq_params import PackedSeqParams diff --git a/tests/unit_tests/models/mimo/test_mimo_grad_sync.py b/tests/unit_tests/models/mimo/test_mimo_grad_sync.py index 33eaa88e907..81f1be17e13 100644 --- a/tests/unit_tests/models/mimo/test_mimo_grad_sync.py +++ b/tests/unit_tests/models/mimo/test_mimo_grad_sync.py @@ -16,9 +16,11 @@ from examples.mimo.training.grad_sync import ( _vision_participation_count, + configure_grad_sync, mark_modality_participation, reset_modality_participation, ) +from megatron.core.models.mimo.config.role import MIMO_LANGUAGE_MODULE_KEY from tests.unit_tests.models.mimo.test_mimo_1f1b_schedule import ( create_hypercomm_grid, destroy_all_grids, @@ -26,6 +28,32 @@ from tests.unit_tests.test_utilities import Utils +def test_configure_grad_sync_installs_production_overlap_hooks(): + def no_sync(): + return None + + def start_grad_sync(*_unused): + return None + + config = SimpleNamespace(no_sync_func=None, grad_sync_func=None) + model = SimpleNamespace( + config=config, + no_sync=no_sync, + start_grad_sync=start_grad_sync, + language_model=None, + modality_submodules={}, + ) + language_grid = SimpleNamespace(get_rank_enum=lambda _name: [[0]]) + topology = SimpleNamespace(grids={MIMO_LANGUAGE_MODULE_KEY: language_grid}, module_pgs={}) + + configure_grad_sync( + SimpleNamespace(overlap_grad_reduce=True, align_grad_reduce=True), model, topology + ) + + assert config.no_sync_func is no_sync + assert config.grad_sync_func is start_grad_sync + + class TestVisionParticipation: @classmethod def setup_class(cls): diff --git a/tests/unit_tests/models/mimo/test_mimo_hetero_grid_args.py b/tests/unit_tests/models/mimo/test_mimo_hetero_grid_args.py index 7429e87434e..fb923c2b37d 100644 --- a/tests/unit_tests/models/mimo/test_mimo_hetero_grid_args.py +++ b/tests/unit_tests/models/mimo/test_mimo_hetero_grid_args.py @@ -103,12 +103,31 @@ def test_llm_cp_must_be_one(): validate_hetero_grid_args(args, WORLD_SIZE_8) +def test_encoder_overlap_requires_grad_reduce(): + args = _layout_8gpu_20l(encoder_ddp_overlap=True, overlap_grad_reduce=False) + with pytest.raises(ValueError, match="requires --overlap-grad-reduce"): + validate_hetero_grid_args(args, WORLD_SIZE_8) + + +def test_encoder_overlap_accepts_uniform_participation_opt_in(): + args = _layout_8gpu_20l(encoder_ddp_overlap=True, overlap_grad_reduce=True) + assert validate_hetero_grid_args(args, WORLD_SIZE_8) == (4, 4) + + def test_llm_only_requires_offset_zero(): args = _layout_8gpu_20l(llm_only=True, llm_offset=4) with pytest.raises(ValueError, match="--llm-only requires --llm-offset 0"): validate_hetero_grid_args(args, WORLD_SIZE_8) +def test_llm_only_rejects_encoder_overlap(): + args = _layout_8gpu_20l( + llm_only=True, llm_offset=0, llm_ep=2, encoder_ddp_overlap=True, overlap_grad_reduce=True + ) + with pytest.raises(ValueError, match="cannot be used with --llm-only"): + validate_hetero_grid_args(args, 4) + + def test_llm_only_covers_world(): # llm tp2/pp1/dp2 = 4 ranks at offset 0; world_size 4 -> covers exactly, no encoder spec. args = _layout_8gpu_20l(llm_only=True, llm_offset=0, llm_ep=2, num_experts=128) diff --git a/tests/unit_tests/models/mimo/test_mimo_hetero_runtime.py b/tests/unit_tests/models/mimo/test_mimo_hetero_runtime.py index 62583aaad82..34c67e07263 100644 --- a/tests/unit_tests/models/mimo/test_mimo_hetero_runtime.py +++ b/tests/unit_tests/models/mimo/test_mimo_hetero_runtime.py @@ -8,7 +8,11 @@ import pytest import torch -from examples.mimo.training.runtime import configure_module_rng, wrap_active_modules_with_ddp +from examples.mimo.training.runtime import ( + _ddp_config_from_args, + configure_module_rng, + wrap_active_modules_with_ddp, +) from examples.mimo.training.topology import ModuleGridSpec, create_topology from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig from megatron.core.enums import ModelType @@ -65,6 +69,43 @@ def _build_unwrapped_mimo_model(topo, bf16=False): return mimo_model +def test_ddp_overlap_config_is_selected_per_module_role(): + args = _args(overlap_grad_reduce=True, overlap_param_gather=True) + + enabled = _ddp_config_from_args(args, enable_overlap=True) + disabled = _ddp_config_from_args(args, enable_overlap=False) + + assert enabled.overlap_grad_reduce + assert enabled.overlap_param_gather + assert not disabled.overlap_grad_reduce + assert not disabled.overlap_param_gather + + +def test_encoder_overlap_opt_in_reaches_encoder_ddp_config(mocker): + encoder = mocker.MagicMock() + wrapped_encoder = mocker.MagicMock() + mimo_model = SimpleNamespace(language_model=None, modality_submodules={ENCODER: encoder}) + topology = SimpleNamespace(module_pgs={ENCODER: mocker.MagicMock()}) + prepare = mocker.patch( + "examples.mimo.training.runtime.prepare_existing_model_chunks_for_distributed_training", + return_value=[wrapped_encoder], + ) + mocker.patch("examples.mimo.training.runtime._freeze_modality_submodule") + mocker.patch("examples.mimo.training.runtime._module_config", return_value=mocker.MagicMock()) + mocker.patch("examples.mimo.training.runtime.print_rank_0") + + wrap_active_modules_with_ddp( + _args(encoder_ddp_overlap=True, overlap_grad_reduce=True, overlap_param_gather=True), + mimo_model, + topology, + ) + + ddp_config = prepare.call_args.kwargs["ddp_config"] + assert ddp_config.overlap_grad_reduce + assert ddp_config.overlap_param_gather + assert mimo_model.modality_submodules[ENCODER] is wrapped_encoder + + def _eight_gpu_topology(): """Encoder dp=4 at ranks 0-3; language dp=4 at ranks 4-7 (non-colocated, tiles world).""" return create_topology( diff --git a/tests/unit_tests/models/mimo/test_mimo_mock_data.py b/tests/unit_tests/models/mimo/test_mimo_mock_data.py index a227a329b4f..6ff70d318ac 100644 --- a/tests/unit_tests/models/mimo/test_mimo_mock_data.py +++ b/tests/unit_tests/models/mimo/test_mimo_mock_data.py @@ -105,6 +105,7 @@ def test_data_adapter_builds_independent_role_specific_loaders(adapter): _args(), _topology(language_rank=True) ) assert all(loader.batch_size == 2 for loader in language_loaders) + assert all(loader.pin_memory for loader in language_loaders) assert len({id(loader.dataset) for loader in language_loaders}) == 3 assert len({loader.dataset.seed for loader in language_loaders}) == 3 language_batch = next(iter(language_loaders[0])) @@ -115,6 +116,7 @@ def test_data_adapter_builds_independent_role_specific_loaders(adapter): _args(), _topology(encoder_rank=True, language_rank=False) ) assert all(loader.batch_size == 4 for loader in encoder_loaders) + assert all(loader.pin_memory for loader in encoder_loaders) encoder_batch = next(iter(encoder_loaders[0])) assert encoder_batch["input_ids"].shape == (4, 8) encoder_inputs = encoder_batch["modality_inputs"][RADIO_ENCODER_MODULE_NAME][ diff --git a/tests/unit_tests/models/mimo/test_mimo_optimizer_consensus.py b/tests/unit_tests/models/mimo/test_mimo_optimizer_consensus.py new file mode 100644 index 00000000000..1bba477ea7f --- /dev/null +++ b/tests/unit_tests/models/mimo/test_mimo_optimizer_consensus.py @@ -0,0 +1,24 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""Distributed test for MimoOptimizer cross-grid step-success consensus.""" + +import pytest +import torch + +from megatron.core.models.mimo.optimizer import MimoOptimizer +from megatron.core.optimizer.optimizer_config import OptimizerConfig +from tests.unit_tests.test_utilities import Utils + + +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="Requires >= 2 ranks.") +def test_step_success_is_world_min(): + """One rank's failed update must propagate to every rank via the MIN reduction.""" + Utils.initialize_distributed() + try: + opt = MimoOptimizer(module_infos={}, config=OptimizerConfig(log_num_zeros_in_grad=False)) + last_rank = torch.distributed.get_world_size() - 1 + opt.step_with_ready_grads = lambda: torch.distributed.get_rank() != last_rank + success, _, _ = opt.step() + assert success is False + finally: + Utils.destroy_model_parallel() diff --git a/tests/unit_tests/models/mimo/test_mimo_overlap_lifecycle.py b/tests/unit_tests/models/mimo/test_mimo_overlap_lifecycle.py new file mode 100644 index 00000000000..b06d1543380 --- /dev/null +++ b/tests/unit_tests/models/mimo/test_mimo_overlap_lifecycle.py @@ -0,0 +1,85 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +"""CPU-only tests for MIMO's nested DDP overlap lifecycle.""" + +from contextlib import contextmanager +from types import SimpleNamespace +from unittest.mock import MagicMock + +from megatron.core.models.mimo.model.base import MimoModel +from megatron.core.models.mimo.optimizer import MimoOptimizer + + +def _overlap_stub(modules): + stub = SimpleNamespace() + stub._active_ddp_modules = lambda: iter(modules) + for name in ( + "no_sync", + "enable_forward_pre_hook", + "disable_forward_pre_hook", + "start_param_sync", + "start_grad_sync", + "free_overlap_buffers", + ): + setattr(stub, name, getattr(MimoModel, name).__get__(stub)) + return stub + + +def _ddp(*, grad_overlap, param_overlap, events, name): + module = MagicMock() + module.ddp_config = SimpleNamespace( + overlap_grad_reduce=grad_overlap, overlap_param_gather=param_overlap + ) + + @contextmanager + def no_sync(): + events.append(f"{name}:enter") + try: + yield + finally: + events.append(f"{name}:exit") + + module.no_sync.side_effect = no_sync + return module + + +def test_nested_overlap_lifecycle_routes_only_to_enabled_modules(): + events = [] + language = _ddp(grad_overlap=True, param_overlap=True, events=events, name="language") + encoder = _ddp(grad_overlap=True, param_overlap=False, events=events, name="encoder") + inactive = _ddp(grad_overlap=False, param_overlap=False, events=events, name="inactive") + model = _overlap_stub([language, encoder, inactive]) + + with model.no_sync(): + events.append("body") + + assert events == ["language:enter", "encoder:enter", "body", "encoder:exit", "language:exit"] + inactive.no_sync.assert_not_called() + + model.enable_forward_pre_hook() + model.disable_forward_pre_hook(param_sync=False) + model.start_param_sync(force_sync=True, force_dispatch=True) + model.start_grad_sync() + model.free_overlap_buffers() + + language.enable_forward_pre_hook.assert_called_once_with() + language.disable_forward_pre_hook.assert_called_once_with(param_sync=False) + language.start_param_sync.assert_called_once_with(force_sync=True, force_dispatch=True) + language.start_grad_sync.assert_called_once_with() + language.free_overlap_buffers.assert_called_once_with() + + encoder.enable_forward_pre_hook.assert_not_called() + encoder.start_param_sync.assert_not_called() + encoder.start_grad_sync.assert_called_once_with() + inactive.start_grad_sync.assert_not_called() + + +def test_mimo_optimizer_stages_each_active_optimizer_before_param_sync(): + language_optimizer = MagicMock() + encoder_optimizer = MagicMock() + optimizer = SimpleNamespace(_active_optimizers=[language_optimizer, encoder_optimizer]) + + MimoOptimizer.prepare_model_params_for_param_sync(optimizer) + + language_optimizer.prepare_model_params_for_param_sync.assert_called_once_with() + encoder_optimizer.prepare_model_params_for_param_sync.assert_called_once_with() diff --git a/tests/unit_tests/models/test_bert_model.py b/tests/unit_tests/models/test_bert_model.py index fb3385b8723..e878978f64e 100644 --- a/tests/unit_tests/models/test_bert_model.py +++ b/tests/unit_tests/models/test_bert_model.py @@ -11,6 +11,7 @@ get_bert_layer_with_transformer_engine_spec, get_bert_layer_with_transformer_engine_submodules, ) +from megatron.core.models.bert.bert_lm_head import BertLMHead from megatron.core.models.bert.bert_model import BertModel from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.enums import AttnBackend, AttnMaskType @@ -92,6 +93,61 @@ def test_post_process_forward(self): assert logits[0].shape[1] == sequence_length assert logits[0].shape[2] == self.bert_model.vocab_size + @pytest.mark.internal + def test_apply_lm_head_default_creates_bert_lm_head(self): + assert isinstance(self.bert_model.lm_head, BertLMHead) + + @pytest.mark.internal + def test_output_layer_bias_false_disables_bias(self): + bert_model = BertModel( + config=self.bert_model.config, + num_tokentypes=0, + transformer_layer_spec=get_bert_layer_with_transformer_engine_spec(), + vocab_size=100, + max_sequence_length=self.bert_model.max_sequence_length, + apply_lm_head=False, + output_layer_bias=False, + ) + + assert bert_model.output_layer.bias is None + + @pytest.mark.internal + def test_apply_lm_head_false_bypasses_head(self): + config: TransformerConfig = self.bert_model.config + sequence_length = self.bert_model.max_sequence_length + micro_batch_size = 2 + + bert_model = BertModel( + config=config, + num_tokentypes=0, + transformer_layer_spec=get_bert_layer_with_transformer_engine_spec(), + vocab_size=100, + max_sequence_length=sequence_length, + apply_lm_head=False, + ) + assert bert_model.lm_head is None + bert_model.cuda() + + encoder_output = {} + bert_model.encoder.register_forward_hook( + lambda module, args, output: encoder_output.setdefault('hidden_states', output) + ) + + data = list(range(sequence_length)) + input_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + attention_mask = torch.ones((micro_batch_size, sequence_length), dtype=bool).cuda() + + logits = bert_model.forward(input_ids=input_ids, attention_mask=attention_mask) + + assert logits[0].shape[0] == micro_batch_size + assert logits[0].shape[1] == sequence_length + assert logits[0].shape[2] == bert_model.vocab_size + + # With apply_lm_head=False, the output_layer must be applied directly to the + # encoder's hidden states, without BertLMHead's dense+GeLU+LayerNorm transform. + expected_logits, _ = bert_model.output_layer(encoder_output['hidden_states']) + torch.testing.assert_close(logits[0], expected_logits.transpose(0, 1).contiguous()) + @pytest.mark.internal def test_qk_layernorm_submodules_are_none(self): # The TE BERT spec leaves q_layernorm/k_layernorm unset (None) instead of hardcoding diff --git a/tests/unit_tests/models/test_gpt_model.py b/tests/unit_tests/models/test_gpt_model.py index 719b3781394..302be05218a 100644 --- a/tests/unit_tests/models/test_gpt_model.py +++ b/tests/unit_tests/models/test_gpt_model.py @@ -1,6 +1,7 @@ # Copyright (c) 2023, NVIDIA CORPORATION. All rights reserved. import inspect +import logging import os from datetime import timedelta from unittest.mock import MagicMock, patch @@ -47,18 +48,29 @@ def setup_method(self, method): use_cpu_initialization=True, embedding_init_method_std=1.0, # Test that we can initialize the embedding weights to something else. ) - self.gpt_model = GPTModel( - config=transformer_config, - transformer_layer_spec=get_gpt_layer_with_transformer_engine_spec(), - vocab_size=100, - max_sequence_length=4, - ) + with patch('megatron.core.models.gpt.gpt_model.log_single_rank') as mock_log_single_rank: + self.gpt_model = GPTModel( + config=transformer_config, + transformer_layer_spec=get_gpt_layer_with_transformer_engine_spec(), + vocab_size=100, + max_sequence_length=4, + ) + self.mock_log_single_rank = mock_log_single_rank def teardown_method(self, method): Utils.destroy_model_parallel() @pytest.mark.internal def test_constructor(self): + self.mock_log_single_rank.assert_called_once() + _, level, message = self.mock_log_single_rank.call_args.args + assert level == logging.WARNING + assert message == ( + "GPTModel IS DEPRECATED. GPTModel is only accepting critical bug fixes, no new " + "features. Please reference the migration guide " + "`docs/user-guide/hybrid-model-migration.md` for details on how to use `HybridModel`" + ) + assert isinstance(self.gpt_model, GPTModel) assert self.gpt_model.max_sequence_length == 4 diff --git a/tests/unit_tests/models/test_gpt_model_batch_invariant.py b/tests/unit_tests/models/test_gpt_model_batch_invariant.py index 9ab7e445c0d..b52fd64f592 100644 --- a/tests/unit_tests/models/test_gpt_model_batch_invariant.py +++ b/tests/unit_tests/models/test_gpt_model_batch_invariant.py @@ -36,6 +36,19 @@ except ImportError: HAVE_FA3 = False +try: + # Blackwell (e.g. GB200) ships FlashAttention-4 instead of FA3; the batch-invariant + # attention paths honor config.flash_attention_version, so these tests run there too. + from flash_attn.cute import flash_attn_varlen_func as _fa4_varlen_func # noqa: F401 + + HAVE_FA4 = True +except ImportError: + HAVE_FA4 = False + +# Batch-invariant mode requires an explicit FlashAttention version; pick the newest +# one available so training and inference run the same kernel. +_BIK_FA_VERSION = 4 if HAVE_FA4 else 3 + class DummyTokenizer: def __init__(self, vocab_size: int, bos: int | None = None, eod: int = 0, pad: int = 0): @@ -86,6 +99,7 @@ def _build_flash_attn_bik_model(seq_len: int, vocab_size: int, hidden_size: int hidden_dropout=0.0, attention_dropout=0.0, batch_invariant_mode=True, + flash_attention_version=_BIK_FA_VERSION, normalization="RMSNorm", params_dtype=torch.bfloat16, attention_backend=AttnBackend.flash, @@ -112,14 +126,21 @@ def _train_forward_logprobs(model: torch.nn.Module, tokens: torch.Tensor) -> tor batch_size, 1, seq_len, seq_len, dtype=torch.bool, device=tokens.device ) with torch.no_grad(): - logits = model(input_ids=tokens, position_ids=position_ids, attention_mask=attention_mask) + # runtime_gather_output matches rl_utils.get_logprobs; without it the model + # asserts once it has served inference requests (in-inference-mode postprocess). + logits = model( + input_ids=tokens, + position_ids=position_ids, + attention_mask=attention_mask, + runtime_gather_output=True, + ) logprobs = selective_log_softmax(logits[:, :-1, :], tokens[:, 1:]) return logprobs @pytest.mark.skipif( - not (is_te_min_version("2.10.0") and HAVE_FA3), - reason="TestGPTModelBatchInvariant requires TE >= 2.10.0 and FlashAttention-3", + not (is_te_min_version("2.10.0") and (HAVE_FA3 or HAVE_FA4)), + reason="TestGPTModelBatchInvariant requires TE >= 2.10.0 and FlashAttention-3 or -4", ) class TestGPTModelBatchInvariant: """End-to-end batch-invariance tests for GPT.""" @@ -199,9 +220,7 @@ def test_dynamic_engine_matches_batched_forward_rl(self): wrapper = GPTInferenceWrapper(inference_model, ctx) tokenizer = DummyTokenizer(vocab_size=vocab_size, bos=None, eod=vocab_size - 1, pad=0) controller = TextGenerationController(wrapper, tokenizer) - engine = DynamicInferenceEngine( - controller=controller, context=ctx, enable_cuda_graph=False, random_seed=123 - ) + engine = DynamicInferenceEngine(controller=controller, context=ctx) base_vals = [3, 15, 27, 39] lengths = [18, 11, 23, 13] @@ -262,7 +281,7 @@ def test_dynamic_engine_is_batch_invariant(self): def _run_engine_with_order(order): ctx = DynamicInferenceContext( - model_config=based_model.config, + model_config=base_model.config, inference_config=InferenceConfig( max_sequence_length=seq_len, buffer_size_gb=0.125, @@ -277,9 +296,7 @@ def _run_engine_with_order(order): wrapper = GPTInferenceWrapper(inference_model, ctx) tokenizer = DummyTokenizer(vocab_size=vocab_size, bos=None, eod=vocab_size - 1, pad=0) controller = TextGenerationController(wrapper, tokenizer) - engine = DynamicInferenceEngine( - controller=controller, context=ctx, enable_cuda_graph=False, random_seed=123 - ) + engine = DynamicInferenceEngine(controller=controller, context=ctx) base_vals = [3, 15, 27, 39] lengths = [18, 11, 23, 13] diff --git a/tests/unit_tests/models/test_hybrid_model.py b/tests/unit_tests/models/test_hybrid_model.py index 6ac6525f30f..b6b23fc6d89 100644 --- a/tests/unit_tests/models/test_hybrid_model.py +++ b/tests/unit_tests/models/test_hybrid_model.py @@ -1,9 +1,12 @@ # Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. +import dataclasses +import functools import os from datetime import timedelta from itertools import accumulate from types import SimpleNamespace +from unittest.mock import patch import pytest import torch @@ -27,13 +30,76 @@ from megatron.core.models.hybrid.hybrid_model import HybridModel, _hybrid_logging_pg_kwargs from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed -from megatron.core.transformer import TransformerConfig +from megatron.core.transformer import MLATransformerConfig, TransformerConfig from megatron.core.transformer.enums import AttnBackend from megatron.core.transformer.module import Float16Module, MegatronModule from megatron.core.transformer.spec_utils import ModuleSpec from megatron.core.utils import divide, is_fa_min_version, is_torch_min_version from tests.unit_tests.test_utilities import Utils +try: + from fast_hadamard_transform import hadamard_transform as _hadamard_transform + + _HAVE_HADAMARD = True +except ImportError: + _HAVE_HADAMARD = False + _hadamard_transform = None + + +def _mock_hadamard_transform(x: torch.Tensor, scale: float = 1.0) -> torch.Tensor: + """Identity-with-scale stand-in for `fast_hadamard_transform.hadamard_transform`. + + Mirrors the helper in `tests/unit_tests/transformer/experimental_attention_variant/ + test_attention_variant_dsa.py` so that DSA forward tests run in containers that + don't ship the upstream library. + """ + return x * scale + + +def _is_dataclass_instance(value): + return dataclasses.is_dataclass(value) and not isinstance(value, type) + + +def _assert_equal_with_partial_contents(left, right, path="root"): + """Assert recursive equality while comparing `partial` objects structurally.""" + if isinstance(left, functools.partial) or isinstance(right, functools.partial): + assert isinstance(left, functools.partial), f"{path}: left is not `partial`" + assert isinstance(right, functools.partial), f"{path}: right is not `partial`" + _assert_equal_with_partial_contents(left.func, right.func, f"{path}.func") + _assert_equal_with_partial_contents(left.args, right.args, f"{path}.args") + _assert_equal_with_partial_contents( + left.keywords or {}, right.keywords or {}, f"{path}.keywords" + ) + return + + if _is_dataclass_instance(left) or _is_dataclass_instance(right): + assert _is_dataclass_instance(left), f"{path}: left is not a dataclass" + assert _is_dataclass_instance(right), f"{path}: right is not a dataclass" + assert type(left) is type(right), f"{path}: dataclass types differ" + for field in dataclasses.fields(left): + if field.compare: + _assert_equal_with_partial_contents( + getattr(left, field.name), getattr(right, field.name), f"{path}.{field.name}" + ) + return + + if isinstance(left, dict) or isinstance(right, dict): + assert isinstance(left, dict), f"{path}: left is not a dict" + assert isinstance(right, dict), f"{path}: right is not a dict" + assert left.keys() == right.keys(), f"{path}: dict keys differ" + for key in left: + _assert_equal_with_partial_contents(left[key], right[key], f"{path}[{key!r}]") + return + + if isinstance(left, (list, tuple)) or isinstance(right, (list, tuple)): + assert type(left) is type(right), f"{path}: sequence types differ" + assert len(left) == len(right), f"{path}: sequence lengths differ" + for index, (left_item, right_item) in enumerate(zip(left, right)): + _assert_equal_with_partial_contents(left_item, right_item, f"{path}[{index}]") + return + + assert left == right, f"{path}: values differ" + class _DummyHybridLayer(MegatronModule): """Minimal same-shape layer used to test HybridModel/mHC plumbing.""" @@ -543,6 +609,12 @@ def test_layer_numbers(self): class TestHybridQKLayernorm: + # Subclasses override these to retarget the same tests at MLA's + # `mla_layer.kv_layernorm` or DSA's `dsa_layer.kv_layernorm`. The base class + # exercises the SelfAttention path with `attention_layer.k_layernorm`. + _attention_layer_attr = 'attention_layer' + _k_norm_attr = 'k_layernorm' + def setup_method(self, method): Utils.initialize_model_parallel(1, 1) model_parallel_cuda_manual_seed(123) @@ -550,7 +622,9 @@ def setup_method(self, method): def teardown_method(self, method): Utils.destroy_model_parallel() - def _build_model(self, **config_overrides): + def _build_model(self, spec=None, **config_overrides): + if spec is None: + spec = hybrid_stack_spec config = TransformerConfig( num_layers=3, hidden_size=256, @@ -560,26 +634,32 @@ def _build_model(self, **config_overrides): ) return HybridModel( config=config, - hybrid_stack_spec=hybrid_stack_spec, + hybrid_stack_spec=spec, vocab_size=100, max_sequence_length=4, hybrid_layer_pattern="M*-", ) def _get_attention_layer(self, model): - """Return the SelfAttention submodule from the attention layer.""" + """Return the self-attention submodule that owns a `q_layernorm`.""" for layer in model.decoder.layers: if hasattr(layer, 'self_attention') and hasattr(layer.self_attention, 'q_layernorm'): return layer.self_attention return None - def test_no_qk_norm_by_default(self): - """Without qk_layernorm, attention has no q/k layernorm.""" + def _get_k_norm(self, attn): + return getattr(attn, self._k_norm_attr) + + def test_trivial_qk_norm_by_default(self): + """Without qk_layernorm, attention has trivial q/k layernorm.""" + from megatron.core.transformer.identity_op import IdentityOp + model = self._build_model() attn = self._get_attention_layer(model) assert attn is not None - assert attn.q_layernorm is None - assert attn.k_layernorm is None + assert attn.q_layernorm is None or isinstance(attn.q_layernorm, IdentityOp) + k_norm = self._get_k_norm(attn) + assert k_norm is None or isinstance(k_norm, IdentityOp) def test_qk_layernorm_from_config(self): """config.qk_layernorm=True creates q/k layernorm even with static spec.""" @@ -589,7 +669,7 @@ def test_qk_layernorm_from_config(self): # TENorm is a factory (__new__ returns a TE LayerNorm/RMSNorm), so we # verify the norm was created rather than checking for a specific type. assert attn.q_layernorm is not None - assert attn.k_layernorm is not None + assert self._get_k_norm(attn) is not None def test_qk_l2_norm_from_config(self): """config.qk_l2_norm=True creates L2Norm q/k layernorm.""" @@ -599,57 +679,649 @@ def test_qk_l2_norm_from_config(self): attn = self._get_attention_layer(model) assert attn is not None assert isinstance(attn.q_layernorm, L2Norm) - assert isinstance(attn.k_layernorm, L2Norm) + assert isinstance(self._get_k_norm(attn), L2Norm) def test_spec_provided_norm_not_overwritten(self): """When the spec already provides q/k layernorm, config doesn't override it.""" import copy - from megatron.core.extensions.transformer_engine import ( - TEDotProductAttention, - TELayerNormColumnParallelLinear, - TERowParallelLinear, - ) - from megatron.core.transformer.attention import SelfAttention, SelfAttentionSubmodules - from megatron.core.transformer.enums import AttnMaskType from megatron.core.transformer.identity_op import IdentityOp - from megatron.core.transformer.spec_utils import ModuleSpec - from megatron.core.transformer.transformer_layer import ( - TransformerLayer, - TransformerLayerSubmodules, - ) - # Build a spec that explicitly sets q/k layernorm to IdentityOp + # Build a spec that explicitly sets q/k layernorm to IdentityOp on the + # attention layer that this subclass exercises. spec = copy.deepcopy(hybrid_stack_spec) - spec.submodules.attention_layer.submodules.self_attention.submodules.q_layernorm = ( - IdentityOp + attn_submodules = getattr( + spec.submodules, self._attention_layer_attr + ).submodules.self_attention.submodules + attn_submodules.q_layernorm = IdentityOp + setattr(attn_submodules, self._k_norm_attr, IdentityOp) + + model = self._build_model(spec=spec, qk_layernorm=True) + attn = self._get_attention_layer(model) + assert attn is not None + assert isinstance(attn.q_layernorm, IdentityOp) + assert isinstance(self._get_k_norm(attn), IdentityOp) + + def test_forward_with_qk_layernorm(self): + """HybridModel forward pass works with qk_layernorm enabled.""" + model = self._build_model(qk_layernorm=True) + model.cuda() + + sequence_length = 4 + micro_batch_size = 2 + data = list(range(sequence_length)) + input_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + position_ids = torch.tensor(data, dtype=torch.int64).repeat((micro_batch_size, 1)).cuda() + attention_mask = torch.ones( + (micro_batch_size, 1, sequence_length, sequence_length), dtype=bool + ).cuda() + + logits = model.forward( + input_ids=input_ids, position_ids=position_ids, attention_mask=attention_mask + ) + + assert logits.shape[0] == micro_batch_size + assert logits.shape[1] == sequence_length + assert logits.shape[2] == 100 + + +class TestHybridMLAQKLayernorm(TestHybridQKLayernorm): + """Tests QK norm configuration of HybridModel with MLA.""" + + _attention_layer_attr = 'mla_layer' + _k_norm_attr = 'kv_layernorm' + + def _build_model(self, spec=None, **config_overrides): + if spec is None: + spec = hybrid_stack_spec + config = MLATransformerConfig( + num_layers=3, + hidden_size=256, + num_attention_heads=4, + use_cpu_initialization=True, + **config_overrides, ) - spec.submodules.attention_layer.submodules.self_attention.submodules.k_layernorm = ( - IdentityOp + return HybridModel( + config=config, + hybrid_stack_spec=spec, + vocab_size=100, + max_sequence_length=4, + hybrid_layer_pattern="M+-", ) - config = TransformerConfig( + def test_qk_l2_norm_from_config(self): + with pytest.raises(ValueError, match="qk_l2_norm is not supported"): + super().test_qk_l2_norm_from_config() + + +class TestHybridDSAQKLayernorm(TestHybridQKLayernorm): + """Tests QK norm configuration of HybridModel with DSA.""" + + _attention_layer_attr = 'dsa_layer' + _k_norm_attr = 'kv_layernorm' + + @pytest.fixture(autouse=True) + def _patch_hadamard_if_needed(self): + if not _HAVE_HADAMARD: + with patch( + 'megatron.core.transformer.experimental_attention_variant.dsa.hadamard_transform', + _mock_hadamard_transform, + ): + yield + else: + yield + + def test_spec_provided_norm_not_overwritten(self): + # DSA cannot fuse the QK norm into the up-projection, so a trivial + # `IdentityOp` spec is auto-promoted to `TENorm` when `qk_layernorm=True`. + # Finer-grained spec-respect behavior is covered by TestDSAQKNormResolution. + pytest.skip("DSA auto-promotes IdentityOp to TENorm; covered by TestDSAQKNormResolution.") + + def _build_model(self, spec=None, **config_overrides): + if spec is None: + spec = hybrid_stack_spec + config_kwargs = dict( num_layers=3, hidden_size=256, num_attention_heads=4, use_cpu_initialization=True, - qk_layernorm=True, + add_bias_linear=False, + # AbsorbedMLASelfAttention forwards `x` and `qr` to the DSA core attention; without + # this, the DSA core attention's forward fails on missing positional arguments. + experimental_attention_variant="dsa", + # DSA-specific settings; defaults are None and DSAIndexer requires them. + dsa_indexer_n_heads=8, + dsa_indexer_head_dim=64, + dsa_indexer_topk=32, + # The indexer-loss path runs in training mode and multiplies by this coefficient; + # leaving it at the default `None` raises `TypeError: ... 'Tensor' and 'NoneType'`. + dsa_indexer_loss_coeff=1.0, + # DSA's `rotate_activation` (Hadamard rotation) only supports bf16 input. + bf16=True, + params_dtype=torch.bfloat16, ) - model = HybridModel( + config_kwargs.update(config_overrides) + config = MLATransformerConfig(**config_kwargs) + return HybridModel( config=config, hybrid_stack_spec=spec, vocab_size=100, max_sequence_length=4, - hybrid_layer_pattern="M*-", + hybrid_layer_pattern="MD-", ) - attn = self._get_attention_layer(model) + + def test_qk_l2_norm_from_config(self): + with pytest.raises(ValueError, match="qk_l2_norm is not supported"): + super().test_qk_l2_norm_from_config() + + +class _MLAQKNormTestBase: + """Common machinery for MLA/DSA QK-norm spec tests. + + Subclasses override `experimental_attention_variant` and + `hybrid_layer_pattern` to target the MLA vs. DSA code path. + """ + + experimental_attention_variant = None + hybrid_layer_pattern = "M+-" + mla_layer_attr = "mla_layer" + + def setup_method(self, method): + Utils.initialize_model_parallel(1, 1) + model_parallel_cuda_manual_seed(123) + + def teardown_method(self, method): + Utils.destroy_model_parallel() + + def _make_spec(self, **submodule_overrides): + """Return a copy of `hybrid_stack_spec` with MLA/DSA submodule overrides.""" + import copy + + spec = copy.deepcopy(hybrid_stack_spec) + mla_submodules = getattr( + spec.submodules, self.mla_layer_attr + ).submodules.self_attention.submodules + for key, value in submodule_overrides.items(): + setattr(mla_submodules, key, value) + return spec + + def _build_model(self, spec=None, **config_overrides): + if spec is None: + spec = hybrid_stack_spec + config_kwargs = dict( + num_layers=3, hidden_size=256, num_attention_heads=4, use_cpu_initialization=True + ) + if self.experimental_attention_variant is not None: + config_kwargs["experimental_attention_variant"] = self.experimental_attention_variant + if self.experimental_attention_variant == "dsa": + # Must not be True for DSA. + config_kwargs.setdefault("add_bias_linear", False) + # DSAIndexer requires these; their config defaults are None. + config_kwargs.setdefault("dsa_indexer_n_heads", 8) + config_kwargs.setdefault("dsa_indexer_head_dim", 64) + config_kwargs.setdefault("dsa_indexer_topk", 32) + + config_kwargs.update(config_overrides) + config = MLATransformerConfig(**config_kwargs) + return HybridModel( + config=config, + hybrid_stack_spec=spec, + vocab_size=100, + max_sequence_length=4, + hybrid_layer_pattern=self.hybrid_layer_pattern, + ) + + def _get_mla_attention(self, model): + """Return the attention submodule for the selected MLA variant, or None.""" + if self.experimental_attention_variant == "dsa": + from megatron.core.transformer.experimental_attention_variant.absorbed_mla import ( + AbsorbedMLASelfAttention, + ) + + attention_cls = AbsorbedMLASelfAttention + else: + from megatron.core.transformer.multi_latent_attention import MLASelfAttention + + attention_cls = MLASelfAttention + + for layer in model.decoder.layers: + if hasattr(layer, 'self_attention') and isinstance(layer.self_attention, attention_cls): + return layer.self_attention + return None + + +class TestMLAQKNormSpecValidation(_MLAQKNormTestBase): + """Tests QK norm spec validation in `MLASelfAttention`. + + These errors guard against silently ignoring a configured norm or + double-applying one through a fused norm+linear. + """ + + experimental_attention_variant = None + hybrid_layer_pattern = "M+-" + mla_layer_attr = "mla_layer" + + def test_q_norm_without_q_lora_rank_raises(self): + """When `q_lora_rank is None`, a non-trivial `q_layernorm` would + never be reached and must error out. + """ + from megatron.core.extensions.transformer_engine import TENorm + + spec = self._make_spec(q_layernorm=TENorm) + with pytest.raises(ValueError, match=r"q_lora_rank is None"): + self._build_model(spec=spec, q_lora_rank=None) + + def test_q_norm_without_q_lora_rank_hint_for_non_fused_linear(self): + """Error message hints at fused linear when `linear_q_proj` is non-fused.""" + from megatron.core.extensions.transformer_engine import TENorm + + spec = self._make_spec(q_layernorm=TENorm) + with pytest.raises(ValueError, match=r"fused norm\+linear for"): + self._build_model(spec=spec, q_lora_rank=None) + + def test_fused_linear_q_up_with_q_norm_raises(self): + """Non-trivial `q_layernorm` combined with a fused `linear_q_up_proj` + would apply the norm twice. + """ + from megatron.core.extensions.transformer_engine import ( + TELayerNormColumnParallelLinear, + TENorm, + ) + + spec = self._make_spec(q_layernorm=TENorm, linear_q_up_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"fused norm\+linear"): + self._build_model(spec=spec) + + def test_fused_linear_kv_up_with_kv_norm_raises(self): + """Non-trivial `kv_layernorm` combined with a fused `linear_kv_up_proj` + would apply the norm twice. + """ + from megatron.core.extensions.transformer_engine import ( + TELayerNormColumnParallelLinear, + TENorm, + ) + + spec = self._make_spec( + kv_layernorm=TENorm, linear_kv_up_proj=TELayerNormColumnParallelLinear + ) + with pytest.raises(ValueError, match=r"fused norm\+linear"): + self._build_model(spec=spec) + + +class TestMLAQKNormResolution(_MLAQKNormTestBase): + """Tests `_resolve_qk_norm_config` for MLA. + + Covers fusion auto-selection, spec overrides, and the "disabled"-path + guards that reject fused/explicit norms when `qk_layernorm` is off. + """ + + experimental_attention_variant = None + hybrid_layer_pattern = "M+-" + mla_layer_attr = "mla_layer" + + def test_qk_layernorm_fuses_kv_up_by_default(self): + """With default (trivial) `kv_layernorm`, enabling `qk_layernorm` + auto-selects the fused `TELayerNormColumnParallelLinear` for KV up. + """ + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + from megatron.core.transformer.identity_op import IdentityOp + + model = self._build_model(qk_layernorm=True) + attn = self._get_mla_attention(model) assert attn is not None - assert isinstance(attn.q_layernorm, IdentityOp) - assert isinstance(attn.k_layernorm, IdentityOp) + assert isinstance(attn.linear_kv_up_proj, TELayerNormColumnParallelLinear) + assert isinstance(attn.kv_layernorm, IdentityOp) + + def test_spec_q_norm_disables_q_up_fusion(self): + """A non-trivial `q_layernorm` from the spec must force a non-fused + `linear_q_up_proj` so the norm isn't applied on top of a fused one. + """ + from megatron.core.extensions.transformer_engine import ( + TEColumnParallelLinear, + TELayerNormColumnParallelLinear, + TENorm, + ) + + spec = self._make_spec(q_layernorm=TENorm) + model = self._build_model(spec=spec, qk_layernorm=True) + attn = self._get_mla_attention(model) + assert attn is not None + assert isinstance(attn.linear_q_up_proj, TEColumnParallelLinear) + assert not isinstance(attn.linear_q_up_proj, TELayerNormColumnParallelLinear) + # The spec's norm is actually used; it's not reset to IdentityOp. + assert attn.q_layernorm is not None + from megatron.core.transformer.identity_op import IdentityOp + + assert not isinstance(attn.q_layernorm, IdentityOp) + + def test_spec_kv_norm_disables_kv_up_fusion(self): + """Mirror of `test_spec_q_norm_disables_q_up_fusion` for KV.""" + from megatron.core.extensions.transformer_engine import ( + TEColumnParallelLinear, + TELayerNormColumnParallelLinear, + TENorm, + ) + + spec = self._make_spec(kv_layernorm=TENorm) + model = self._build_model(spec=spec, qk_layernorm=True) + attn = self._get_mla_attention(model) + assert attn is not None + assert isinstance(attn.linear_kv_up_proj, TEColumnParallelLinear) + assert not isinstance(attn.linear_kv_up_proj, TELayerNormColumnParallelLinear) + from megatron.core.transformer.identity_op import IdentityOp + + assert not isinstance(attn.kv_layernorm, IdentityOp) + + def test_disabled_qk_layernorm_rejects_fused_linear_q_up(self): + """When `qk_layernorm` is off, spec must not force fused linear_q_up_proj.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + + spec = self._make_spec(linear_q_up_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"supposed to be disabled"): + self._build_model(spec=spec) + + def test_disabled_qk_layernorm_rejects_fused_linear_kv_up(self): + """When `qk_layernorm` is off, spec must not force fused linear_kv_up_proj.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + + spec = self._make_spec(linear_kv_up_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"supposed to be disabled"): + self._build_model(spec=spec) + + def test_disabled_qk_layernorm_rejects_spec_norms(self): + """When `qk_layernorm` is off, spec must not carry explicit q/kv layernorms.""" + from megatron.core.extensions.transformer_engine import TENorm + + for overrides in ( + {"q_layernorm": TENorm}, + {"kv_layernorm": TENorm}, + {"q_layernorm": TENorm, "kv_layernorm": TENorm}, + ): + spec = self._make_spec(**overrides) + with pytest.raises(ValueError, match=r"supposed to be disabled"): + self._build_model(spec=spec) + + +class TestDSAQKNormResolution(_MLAQKNormTestBase): + """Tests `_resolve_qk_norm_config` for DSA. + + DSA requires non-fused Q/KV up projections and explicit norms; + the fused optimization valid for MLA must be rejected here. + """ + + experimental_attention_variant = "dsa" + hybrid_layer_pattern = "MD-" + mla_layer_attr = "dsa_layer" + + def test_qk_layernorm_uses_unfused_linear_and_te_norm(self): + """With default spec, DSA + `qk_layernorm=True` uses non-fused + `TEColumnParallelLinear` and `TENorm` for Q/KV. + """ + from megatron.core.extensions.transformer_engine import ( + TEColumnParallelLinear, + TELayerNormColumnParallelLinear, + ) + from megatron.core.transformer.identity_op import IdentityOp - def test_forward_with_qk_layernorm(self): - """HybridModel forward pass works with qk_layernorm enabled.""" model = self._build_model(qk_layernorm=True) + attn = self._get_mla_attention(model) + assert attn is not None + assert isinstance(attn.linear_q_up_proj, TEColumnParallelLinear) + assert not isinstance(attn.linear_q_up_proj, TELayerNormColumnParallelLinear) + assert isinstance(attn.linear_kv_up_proj, TEColumnParallelLinear) + assert not isinstance(attn.linear_kv_up_proj, TELayerNormColumnParallelLinear) + assert not isinstance(attn.q_layernorm, IdentityOp) + assert not isinstance(attn.kv_layernorm, IdentityOp) + + def test_qk_layernorm_without_q_lora_rank_raises(self): + """DSA cannot apply Q norm when `q_lora_rank is None`.""" + with pytest.raises(ValueError, match=r"q_lora_rank is None.*not supported for DSA"): + self._build_model(qk_layernorm=True, q_lora_rank=None) + + def test_qk_layernorm_rejects_fused_linear_q_up(self): + """DSA does not support the fused norm+linear optimization.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + + spec = self._make_spec(linear_q_up_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"not supported for DSA"): + self._build_model(spec=spec, qk_layernorm=True) + + def test_qk_layernorm_without_q_lora_rejects_fused_linear_q(self): + """DSA does not support fused `linear_q_proj` when `q_lora_rank=None`.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + + spec = self._make_spec(linear_q_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"not supported for DSA"): + self._build_model(spec=spec, qk_layernorm=True, q_lora_rank=None) + + def test_disabled_qk_layernorm_rejects_fused_linear_kv_up(self): + """When `qk_layernorm` is off, spec must not force fused linear_kv_up_proj.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + + spec = self._make_spec(linear_kv_up_proj=TELayerNormColumnParallelLinear) + with pytest.raises(ValueError, match=r"supposed to be disabled"): + self._build_model(spec=spec) + + def test_disabled_qk_layernorm_rejects_spec_norms(self): + """When `qk_layernorm` is off, spec must not carry explicit q/kv layernorms.""" + from megatron.core.extensions.transformer_engine import TENorm + + for overrides in ( + {"q_layernorm": TENorm}, + {"kv_layernorm": TENorm}, + {"q_layernorm": TENorm, "kv_layernorm": TENorm}, + ): + spec = self._make_spec(**overrides) + with pytest.raises(ValueError, match=r"supposed to be disabled"): + self._build_model(spec=spec) + + +class TestMLADownProjFusion: + """Tests `HybridStack._fuse_mla_down_proj`. + + The method rewrites the MLA `ModuleSpec` in place on a deep-copied + `HybridStackSubmodules` when `config.mla_down_proj_fusion=True`, swapping + the self-attention module to `FusedMLASelfAttention` and collapsing the + separate q/kv down projections into a single fused `linear_qkv_down_proj` + that also absorbs the input layernorm. + """ + + def setup_method(self, method): + Utils.initialize_model_parallel(1, 1) + model_parallel_cuda_manual_seed(123) + + def teardown_method(self, method): + Utils.destroy_model_parallel() + + def _fresh_submodules(self): + """Return a deep copy of `hybrid_stack_spec.submodules` so tests don't + share state through `hybrid_stack_spec`. + """ + import copy + + return copy.deepcopy(hybrid_stack_spec.submodules) + + def _call_fuse(self, submodules, *, mla_down_proj_fusion): + """Invoke `_fuse_mla_down_proj` as an unbound method with a minimal + stub for `self`. The method only reads `self.config`, so we can avoid + constructing a full `HybridStack`. + """ + from megatron.core.models.hybrid.hybrid_block import HybridStack + + stub = SimpleNamespace(config=SimpleNamespace(mla_down_proj_fusion=mla_down_proj_fusion)) + # Mimic the call-site check in `HybridStack.__init__`. + if getattr(stub.config, "mla_down_proj_fusion", False): + submodules = HybridStack._fuse_mla_down_proj(stub, submodules) + return submodules + + def _build_model(self, pattern="M+-", **config_overrides): + config_kwargs = dict( + num_layers=3, hidden_size=256, num_attention_heads=4, use_cpu_initialization=True + ) + config_kwargs.update(config_overrides) + config = MLATransformerConfig(**config_kwargs) + return HybridModel( + config=config, + hybrid_stack_spec=hybrid_stack_spec, + vocab_size=100, + max_sequence_length=4, + hybrid_layer_pattern=pattern, + ) + + def _get_layer_with_mla(self, model): + """Return the layer whose self-attention is an `MLASelfAttention` + (which includes its `FusedMLASelfAttention` subclass). + """ + from megatron.core.transformer.multi_latent_attention import MLASelfAttention + + for layer in model.decoder.layers: + if hasattr(layer, 'self_attention') and isinstance( + layer.self_attention, MLASelfAttention + ): + return layer + return None + + def test_disabled_returns_spec_unchanged(self): + """Flag off: method returns the same object, no copying or rewriting.""" + submodules = self._fresh_submodules() + result = self._call_fuse(submodules, mla_down_proj_fusion=False) + assert result is submodules + + def test_enabled_rewrites_mla_spec(self): + """Flag on: MLA spec is swapped to the fused module and fused linear.""" + from megatron.core.extensions.transformer_engine import TELayerNormColumnParallelLinear + from megatron.core.transformer.identity_op import IdentityOp + from megatron.core.transformer.multi_latent_attention import FusedMLASelfAttention + + submodules = self._fresh_submodules() + result = self._call_fuse(submodules, mla_down_proj_fusion=True) + + mla_spec = result.mla_layer + assert mla_spec.submodules.input_layernorm is IdentityOp + assert mla_spec.submodules.self_attention.module is FusedMLASelfAttention + + attn_submodules = mla_spec.submodules.self_attention.submodules + assert attn_submodules.linear_qkv_down_proj is TELayerNormColumnParallelLinear + assert attn_submodules.linear_q_down_proj is None + assert attn_submodules.linear_kv_down_proj is None + + def test_enabled_sets_sharded_state_dict_keys_map(self): + """The keys map is written on the MLA layer submodules for checkpoint + compatibility with pre-fusion checkpoints. + """ + submodules = self._fresh_submodules() + result = self._call_fuse(submodules, mla_down_proj_fusion=True) + + keys_map = result.mla_layer.submodules.sharded_state_dict_keys_map + assert keys_map == { + "self_attention.linear_q_down_proj.layer_norm_": "input_layernorm.", + "self_attention.linear_kv_down_proj.layer_norm_": "input_layernorm.", + "self_attention.linear_qkv_down_proj.layer_norm_": "input_layernorm.", + } + + def test_enabled_deep_copies_input_submodules(self): + """The caller's submodules object must not be mutated – the method + deep-copies before rewriting, so callers can safely reuse their spec. + """ + from megatron.core.transformer.multi_latent_attention import ( + FusedMLASelfAttention, + MLASelfAttention, + ) + + submodules = self._fresh_submodules() + original_mla_module = submodules.mla_layer.submodules.self_attention.module + original_q_down_proj = ( + submodules.mla_layer.submodules.self_attention.submodules.linear_q_down_proj + ) + assert original_mla_module is MLASelfAttention # sanity check of baseline + + result = self._call_fuse(submodules, mla_down_proj_fusion=True) + + # Original is unchanged. + assert submodules.mla_layer.submodules.self_attention.module is original_mla_module + assert ( + submodules.mla_layer.submodules.self_attention.submodules.linear_q_down_proj + is original_q_down_proj + ) + # And result is a different object than the input. + assert result is not submodules + assert result.mla_layer is not submodules.mla_layer + # Plus the fused module only shows up on the returned copy. + assert result.mla_layer.submodules.self_attention.module is FusedMLASelfAttention + + def test_enabled_leaves_dsa_layer_alone(self): + """MLA fusion must not rewrite the absorbed DSA attention specification.""" + from megatron.core.transformer.experimental_attention_variant.absorbed_mla import ( + AbsorbedMLASelfAttention, + ) + from megatron.core.transformer.multi_latent_attention import FusedMLASelfAttention + + submodules = self._fresh_submodules() + result = self._call_fuse(submodules, mla_down_proj_fusion=True) + + assert result.dsa_layer.submodules.self_attention.module is AbsorbedMLASelfAttention + assert result.dsa_layer.submodules.self_attention.module is not FusedMLASelfAttention + # DSA's down projections must remain non-`None` (they're still used + # via the unfused path). + assert result.dsa_layer.submodules.self_attention.submodules.linear_q_down_proj is not None + assert result.dsa_layer.submodules.self_attention.submodules.linear_kv_down_proj is not None + + def test_enabled_leaves_non_mla_layers_alone(self): + """Unrelated layer specs (mamba, attention, mlp) must survive unchanged.""" + submodules = self._fresh_submodules() + original_mamba = submodules.mamba_layer + original_attention = submodules.attention_layer + original_mlp = submodules.mlp_layer + + result = self._call_fuse(submodules, mla_down_proj_fusion=True) + + _assert_equal_with_partial_contents(result.mamba_layer, original_mamba) + _assert_equal_with_partial_contents(result.attention_layer, original_attention) + _assert_equal_with_partial_contents(result.mlp_layer, original_mlp) + + def test_model_uses_fused_mla_when_enabled(self): + """Integration: a full HybridModel built with the flag uses + `FusedMLASelfAttention`. + """ + from megatron.core.transformer.multi_latent_attention import FusedMLASelfAttention + + model = self._build_model(mla_down_proj_fusion=True) + layer = self._get_layer_with_mla(model) + assert layer is not None + assert isinstance(layer.self_attention, FusedMLASelfAttention) + # And the fused down projection is present on the attention module. + assert hasattr(layer.self_attention, "linear_qkv_down_proj") + + def test_model_uses_unfused_mla_when_disabled(self): + """Integration: with the flag off, MLA layers use the standard + `MLASelfAttention` (never the fused subclass). + """ + from megatron.core.transformer.multi_latent_attention import ( + FusedMLASelfAttention, + MLASelfAttention, + ) + + model = self._build_model(mla_down_proj_fusion=False) + layer = self._get_layer_with_mla(model) + assert layer is not None + assert isinstance(layer.self_attention, MLASelfAttention) + assert not isinstance(layer.self_attention, FusedMLASelfAttention) + + def test_enabled_replaces_input_layernorm_with_identity(self): + """Integration: because the fused down-proj absorbs the input + layernorm, the transformer layer's own `input_layernorm` must be + `IdentityOp`. + """ + from megatron.core.transformer.identity_op import IdentityOp + + model = self._build_model(mla_down_proj_fusion=True) + layer = self._get_layer_with_mla(model) + assert layer is not None + assert isinstance(layer.input_layernorm, IdentityOp) + + def test_forward_with_fused_mla(self): + """Integration: forward pass works with `mla_down_proj_fusion=True`.""" + model = self._build_model(mla_down_proj_fusion=True) model.cuda() sequence_length = 4 diff --git a/tests/unit_tests/models/test_hybrid_moe_model.py b/tests/unit_tests/models/test_hybrid_moe_model.py index b8c1c4b4e56..971ab92fa13 100644 --- a/tests/unit_tests/models/test_hybrid_moe_model.py +++ b/tests/unit_tests/models/test_hybrid_moe_model.py @@ -115,11 +115,14 @@ "experimental_attention_variant": None, "experimental_attention_variant_loss_scale_func": None, "expert_model_parallel_size": 4, + "expert_gtp_weight_remat_size": 1, + "expert_tensor_parallel_num_weight_shards": 1, "expert_tensor_parallel_size": 1, "external_cuda_graph": False, "ffn_hidden_size": 1856, "finalize_model_grads_func": None, "first_last_layers_bf16": False, + "flash_attention_version": None, "flash_decode": False, "fp16": False, "fp32_residual_connection": False, @@ -142,6 +145,7 @@ "fused_residual_rmsnorm": False, "fused_single_qkv_rope": False, "gated_linear_unit": False, + "gtp_weight_remat_size": 1, "glu_linear_offset": 0.0, "grad_scale_func": None, "mtp_grad_scale_func": None, @@ -208,7 +212,7 @@ "moe_layer_recompute": False, "moe_n_hash_layers": 0, "moe_ncclep_static_shape": False, - "moe_ncclep_use_symm_mem": False, + "moe_ncclep_zero_copy": False, "moe_pad_expert_input_to_capacity": False, "moe_pad_experts_for_cuda_graph_inference": False, "moe_paged_stash": False, @@ -229,6 +233,7 @@ "moe_router_padding_for_fp8": False, "moe_router_padding_for_quantization": False, "moe_router_pre_softmax": False, + "moe_router_quantile_balancing_ema": 0.0, "moe_router_score_function": "sigmoid", "moe_router_topk": 6, "moe_router_topk_limited_devices": None, @@ -301,6 +306,7 @@ "softmax_type": "vanilla", "symmetric_ar_type": None, "tensor_model_parallel_size": 2, + "tensor_parallel_num_weight_shards": 2, "test_mode": False, "thd_max_packed_sequences": None, "timers": None, diff --git a/tests/unit_tests/models/test_nemo_audio_processor.py b/tests/unit_tests/models/test_nemo_audio_processor.py new file mode 100644 index 00000000000..8e97e0b8bab --- /dev/null +++ b/tests/unit_tests/models/test_nemo_audio_processor.py @@ -0,0 +1,455 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. +# SPDX-License-Identifier: BSD-3-Clause + +"""Tests for the data-side NeMo audio processor. + +Covers the cumulative-prefix slice primitives +(``num_frames_from_num_samples`` / ``num_embeddings_from_num_samples``), +``slice_range`` waveform cropping, and log-mel materialization. ``audio_ref`` is +duck-typed (a tiny ``SimpleNamespace`` stand-in) — the processor never imports +the data library's ``AudioRef`` type. +""" + +from types import SimpleNamespace + +import pytest + +torch = pytest.importorskip("torch") + +from megatron.core.models.audio import audio_processor +from megatron.core.models.audio.audio_feature_config import ( + NemoAudioFeatureConfig, + NemoTransformerAudioTokenEstimator, +) +from megatron.core.models.audio.audio_processor import NemoAudioProcessor + + +def _audio_ref(**kwargs): + fields = dict( + data=None, + sample_rate=None, + num_samples=None, + slice_range=None, + num_frames=None, + feature_dim=None, + ) + fields.update(kwargs) + return SimpleNamespace(**fields) + + +def _processor(): + return NemoAudioProcessor( + token_estimator=NemoTransformerAudioTokenEstimator( + encoder_time_stride=4, stack_factor=2, pre_encode="conv" + ), + feature_config=NemoAudioFeatureConfig( + sample_rate=16000, window_stride=0.01, n_window_stride=None, dither=0.0 + ), + ) + + +# --------------------------------------------------------------------------- +# Slice primitives +# --------------------------------------------------------------------------- + + +def test_num_frames_from_num_samples_zero(): + assert _processor().num_frames_from_num_samples(0) == 0 + + +def test_num_frames_from_num_samples_one_second(): + # 16000 samples @ hop=160. + assert _processor().num_frames_from_num_samples(16000) == 100 + + +def test_num_embeddings_from_num_samples_one_second(): + # 100 frames -> floor(100 / 4) encoder steps -> ceil(25 / 2) embeddings. + assert _processor().num_embeddings_from_num_samples(16000) == 13 + + +def test_num_embeddings_from_num_samples_zero(): + assert _processor().num_embeddings_from_num_samples(0) == 0 + + +@pytest.mark.parametrize( + "boundaries", + [ + [0, 8000, 16000], # 0.5s, 1.0s halves + [0, 3200, 12800, 28800, 64000, 100000, 160000], # arbitrary 10s split + ], +) +def test_slice_contributions_sum_to_full_embedding_count(boundaries): + """Cumulative-prefix invariant: contributions of disjoint slices sum to the + unsliced total. This is the property that justifies the primitives.""" + p = _processor() + total = p.num_embeddings_from_num_samples(boundaries[-1]) + cum = sum( + p.num_embeddings_from_num_samples(e) - p.num_embeddings_from_num_samples(s) + for s, e in zip(boundaries[:-1], boundaries[1:]) + ) + assert cum == total + + +@pytest.mark.parametrize( + "boundaries", [[0, 8000, 16000], [0, 3200, 12800, 28800, 64000, 100000, 160000]] +) +def test_slice_contributions_sum_to_full_frame_count(boundaries): + p = _processor() + total = p.num_frames_from_num_samples(boundaries[-1]) + cum = sum( + p.num_frames_from_num_samples(e) - p.num_frames_from_num_samples(s) + for s, e in zip(boundaries[:-1], boundaries[1:]) + ) + assert cum == total + + +def test_independent_slice_lengths_drift_from_total(): + """The naive ``f(end - start)`` approach is provably WRONG: ceil divisions + mean per-slice-length counts do not sum to the unsliced total. Use the + subtraction (cumulative-prefix) contract, not addition.""" + p = _processor() + assert p.num_embeddings_from_num_samples(16000) == 13 + assert p.num_embeddings_from_num_samples(8000) == 6 + # 6 + 6 != 13. + assert p.num_embeddings_from_num_samples(8000) + p.num_embeddings_from_num_samples(8000) != 13 + + +# --------------------------------------------------------------------------- +# Waveform normalization / slicing +# --------------------------------------------------------------------------- + + +def test_normalize_waveform_crops_to_slice_range(monkeypatch): + waveform = torch.arange(20, dtype=torch.float32) + + def fake_load_waveform(audio_spec): + del audio_spec + return waveform, 16000 + + monkeypatch.setattr(audio_processor, "_load_waveform_from_spec", fake_load_waveform) + # num_samples is the full source length; slice_range carries the crop window. + audio = _audio_ref( + data={"kind": "avdecoder"}, sample_rate=16000, num_samples=20, slice_range=(5, 9) + ) + + cropped, decoded_sample_rate = audio_processor._normalize_mono_waveform(audio) + + assert decoded_sample_rate == 16000 + assert cropped.tolist() == [5.0, 6.0, 7.0, 8.0] + + +def test_infer_num_samples_uses_slice_range_length(): + # slice_range defines the effective length even though num_samples is the full source. + audio = _audio_ref(data={"kind": "avdecoder"}, num_samples=20, slice_range=(5, 9)) + assert audio_processor._infer_num_samples(audio) == 4 + + +# --------------------------------------------------------------------------- +# Materialization +# --------------------------------------------------------------------------- + + +def test_materialize_returns_time_major_log_mel(): + p = _processor() + waveform = torch.zeros(16000, dtype=torch.float32) + audio = _audio_ref(data=waveform, sample_rate=16000, num_samples=16000) + + log_mel, valid_frames = p.materialize(audio) + + assert valid_frames == 100 + assert log_mel.shape == (100, p.input_feature_dim) + assert log_mel.dtype == torch.float32 + + +def test_materialize_empty_waveform_yields_no_frames(): + p = _processor() + audio = _audio_ref(data=torch.zeros(0, dtype=torch.float32), sample_rate=16000, num_samples=0) + log_mel, valid_frames = p.materialize(audio) + assert valid_frames == 0 + assert log_mel.shape == (0, p.input_feature_dim) + + +# --------------------------------------------------------------------------- +# Public methods / properties +# --------------------------------------------------------------------------- + + +def test_processor_properties(): + p = _processor() + assert p.sample_rate == 16000 + assert p.input_feature_dim == p._n_mels + + +def test_compute_num_frames_and_embeddings_from_waveform_ref(): + p = _processor() + audio = _audio_ref(data=torch.zeros(16000, dtype=torch.float32), sample_rate=16000) + # 16000 samples @ hop=160 -> 100 frames -> 13 embeddings (see slice-primitive tests). + assert p.compute_num_frames(audio) == 100 + assert p.compute_num_embeddings(audio) == 13 + + +def test_validate_sample_rate_mismatch_raises(): + p = _processor() + audio = _audio_ref(data=torch.zeros(16000, dtype=torch.float32), sample_rate=8000) + with pytest.raises(ValueError, match="Expected audio sample rate 16000"): + p.compute_num_frames(audio) + + +# --------------------------------------------------------------------------- +# Lazy AV-decoder decode chain (duck-typed decoder fakes) +# --------------------------------------------------------------------------- + + +class _FakeAVData: + def __init__(self, clips): + self.audio_clips = clips + + +class _FakeDecoder: + """Minimal AVDecoder-like stand-in exposing get_audio + sample-rate probes.""" + + def __init__(self, clips, samples_per_second=None): + self._clips = clips + self._samples_per_second = samples_per_second + + def get_audio(self): + return _FakeAVData(self._clips) + + def get_audio_samples_per_second(self): + return self._samples_per_second + + +def test_audio_clip_to_float32_1d_float_is_unsqueezed(): + out = audio_processor._audio_clip_to_float32(torch.tensor([0.1, 0.2], dtype=torch.float32)) + assert out.shape == (1, 2) + assert out.dtype == torch.float32 + + +def test_audio_clip_to_float32_from_python_list(): + out = audio_processor._audio_clip_to_float32([0.0, 1.0, -1.0]) + assert out.shape == (1, 3) + assert out.dtype == torch.float32 + + +def test_audio_clip_to_float32_uint8_is_centered_and_scaled(): + out = audio_processor._audio_clip_to_float32(torch.tensor([0, 128, 255], dtype=torch.uint8)) + assert torch.allclose(out[0], torch.tensor([-1.0, 0.0, (255 - 128) / 128.0])) + + +def test_audio_clip_to_float32_int16_is_scaled_by_max(): + out = audio_processor._audio_clip_to_float32(torch.tensor([0, 32767], dtype=torch.int16)) + assert torch.allclose(out[0], torch.tensor([0.0, 1.0])) + + +def test_audio_clip_to_float32_rejects_bad_ndim(): + with pytest.raises(ValueError, match="Unsupported decoded audio clip shape"): + audio_processor._audio_clip_to_float32(torch.zeros(2, 2, 2)) + + +def test_audio_clip_to_float32_rejects_bad_dtype(): + with pytest.raises(ValueError, match="Unsupported decoded audio dtype"): + audio_processor._audio_clip_to_float32(torch.tensor([True, False])) + + +def test_decoder_sample_rate_prefers_explicit(): + assert audio_processor._decoder_sample_rate(object(), 22050) == 22050 + + +def test_decoder_sample_rate_from_samples_per_second(): + dec = _FakeDecoder(clips=[], samples_per_second=16000) + assert audio_processor._decoder_sample_rate(dec, None) == 16000 + + +def test_decoder_sample_rate_from_metadata(): + class _MetaDecoder: + def get_metadata(self, **kwargs): + del kwargs + return SimpleNamespace(audio_sample_rate=8000) + + assert audio_processor._decoder_sample_rate(_MetaDecoder(), None) == 8000 + + +def test_decoder_sample_rate_none_when_unavailable(): + assert audio_processor._decoder_sample_rate(object(), None) is None + + +def test_resolve_lazy_media_unwraps_get_and_sequence(): + dec = _FakeDecoder(clips=[]) + lazy = SimpleNamespace(get=lambda: [dec]) + assert audio_processor._resolve_lazy_media(lazy) is dec + + +def test_resolve_lazy_media_empty_sequence_raises(): + with pytest.raises(ValueError, match="empty sequence"): + audio_processor._resolve_lazy_media([]) + + +def test_decode_avdecoder_concatenates_clips(): + dec = _FakeDecoder( + clips=[ + torch.tensor([0.0, 1.0], dtype=torch.float32), + torch.tensor([2.0], dtype=torch.float32), + ], + samples_per_second=16000, + ) + waveform, sr = audio_processor._decode_avdecoder(dec, "") + assert sr == 16000 + assert waveform.shape == (1, 3) + assert waveform[0].tolist() == [0.0, 1.0, 2.0] + + +def test_decode_avdecoder_rejects_non_decoder(): + with pytest.raises(ValueError, match="Expected AVDecoder-like"): + audio_processor._decode_avdecoder(object(), "") + + +def test_decode_avdecoder_rejects_missing_clips(): + with pytest.raises(ValueError, match="did not contain audio clips"): + audio_processor._decode_avdecoder(_FakeDecoder(clips=[]), "") + + +def test_load_waveform_from_spec_dispatches_avdecoder(): + dec = _FakeDecoder(clips=[torch.tensor([0.5], dtype=torch.float32)], samples_per_second=16000) + waveform, sr = audio_processor._load_waveform_from_spec( + {"kind": "avdecoder", "decoder": dec, "sample_rate": 16000} + ) + assert sr == 16000 + assert waveform.shape == (1, 1) + + +def test_load_waveform_from_spec_rejects_unknown_kind(): + with pytest.raises(ValueError, match="Unsupported audio kind"): + audio_processor._load_waveform_from_spec({"kind": "wav"}) + + +# --------------------------------------------------------------------------- +# Sample-rate / tolerance resolution +# --------------------------------------------------------------------------- + + +def test_resolve_sample_rate_prefers_audio_ref(): + audio = _audio_ref(sample_rate=16000, data={}) + assert audio_processor._resolve_sample_rate(audio, 8000) == 16000 + + +def test_resolve_sample_rate_falls_back_to_decoded(): + audio = _audio_ref(sample_rate=None, data={}) + assert audio_processor._resolve_sample_rate(audio, 8000) == 8000 + + +def test_resolve_sample_rate_falls_back_to_data_dict(): + audio = _audio_ref(sample_rate=None, data={"sampling_rate": 22050}) + assert audio_processor._resolve_sample_rate(audio, None) == 22050 + + +def test_resolve_sample_rate_none_when_unknown(): + audio = _audio_ref(sample_rate=None, data=torch.zeros(1)) + assert audio_processor._resolve_sample_rate(audio, None) is None + + +def test_audio_num_sample_tolerance(): + audio = _audio_ref(sample_rate=16000, data={}) + # ceil(0.5 * 16000) = 8000 samples of allowed drift. + assert audio_processor._audio_num_sample_tolerance(audio, None) == 8000 + + +def test_audio_num_sample_tolerance_zero_without_sample_rate(): + audio = _audio_ref(sample_rate=None, data=torch.zeros(1)) + assert audio_processor._audio_num_sample_tolerance(audio, None) == 0 + + +# --------------------------------------------------------------------------- +# Waveform normalization branches +# --------------------------------------------------------------------------- + + +def test_normalize_stereo_is_averaged_to_mono(): + data = torch.tensor([[0.0, 2.0], [2.0, 4.0]], dtype=torch.float32) # [C=2, T=2] + audio = _audio_ref(data=data, sample_rate=16000) + waveform, _ = audio_processor._normalize_mono_waveform(audio) + assert waveform.tolist() == [1.0, 3.0] + + +def test_normalize_pads_up_to_num_samples_within_tolerance(): + audio = _audio_ref(data=torch.ones(10, dtype=torch.float32), sample_rate=16000, num_samples=13) + waveform, _ = audio_processor._normalize_mono_waveform(audio) + assert waveform.shape[0] == 13 + assert waveform[10:].tolist() == [0.0, 0.0, 0.0] + + +def test_normalize_crops_down_to_num_samples(): + audio = _audio_ref(data=torch.arange(10, dtype=torch.float32), sample_rate=16000, num_samples=4) + waveform, _ = audio_processor._normalize_mono_waveform(audio) + assert waveform.tolist() == [0.0, 1.0, 2.0, 3.0] + + +def test_normalize_rejects_num_samples_beyond_tolerance(): + audio = _audio_ref( + data=torch.ones(10, dtype=torch.float32), sample_rate=16000, num_samples=9000 + ) + with pytest.raises(ValueError, match="exceeds waveform length"): + audio_processor._normalize_mono_waveform(audio) + + +def test_normalize_rejects_non_float32(): + audio = _audio_ref(data=torch.zeros(10, dtype=torch.float64), sample_rate=16000) + with pytest.raises(ValueError, match="Expected raw float32 waveform"): + audio_processor._normalize_mono_waveform(audio) + + +def test_normalize_rejects_bad_ndim(): + audio = _audio_ref(data=torch.zeros(2, 2, 2, dtype=torch.float32), sample_rate=16000) + with pytest.raises(ValueError, match="Unsupported waveform shape"): + audio_processor._normalize_mono_waveform(audio) + + +def test_normalize_rejects_unsupported_data_type(): + audio = _audio_ref(data="not-a-waveform", sample_rate=16000) + with pytest.raises(ValueError, match="must be a raw float32 waveform"): + audio_processor._normalize_mono_waveform(audio) + + +def test_normalize_rejects_bad_slice_range(): + audio = _audio_ref( + data=torch.ones(10, dtype=torch.float32), sample_rate=16000, slice_range=(5, 2) + ) + with pytest.raises(ValueError, match="slice_range must satisfy"): + audio_processor._normalize_mono_waveform(audio) + + +# --------------------------------------------------------------------------- +# _infer_num_samples branches +# --------------------------------------------------------------------------- + + +def test_infer_num_samples_from_num_samples_field(): + audio = _audio_ref(num_samples=1234, data=torch.zeros(1, dtype=torch.float32)) + assert audio_processor._infer_num_samples(audio) == 1234 + + +def test_infer_num_samples_from_1d_tensor(): + audio = _audio_ref(data=torch.zeros(500, dtype=torch.float32)) + assert audio_processor._infer_num_samples(audio) == 500 + + +def test_infer_num_samples_from_2d_tensor_uses_time_dim(): + audio = _audio_ref(data=torch.zeros(2, 640, dtype=torch.float32)) + assert audio_processor._infer_num_samples(audio) == 640 + + +def test_infer_num_samples_from_spec(): + dec = _FakeDecoder(clips=[torch.zeros(320, dtype=torch.float32)], samples_per_second=16000) + audio = _audio_ref(data={"kind": "avdecoder", "decoder": dec}) + assert audio_processor._infer_num_samples(audio) == 320 + + +def test_infer_num_samples_rejects_unsupported_data(): + audio = _audio_ref(data=42) + with pytest.raises(ValueError, match="must be a raw float32 waveform"): + audio_processor._infer_num_samples(audio) + + +def test_infer_num_samples_rejects_non_float32(): + audio = _audio_ref(data=torch.zeros(10, dtype=torch.int32)) + with pytest.raises(ValueError, match="Expected raw float32 waveform"): + audio_processor._infer_num_samples(audio) diff --git a/tests/unit_tests/optimizer/test_param_group_identifier_keys.py b/tests/unit_tests/optimizer/test_param_group_identifier_keys.py new file mode 100644 index 00000000000..6c79a54c53c --- /dev/null +++ b/tests/unit_tests/optimizer/test_param_group_identifier_keys.py @@ -0,0 +1,271 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""Tests for ``param_group_identifier_keys`` and the param-group save/load matching. + +The identifier tuple is the fingerprint used by +:meth:`MegatronOptimizer._filter_and_reorder_param_groups` and +:meth:`DistributedOptimizer.load_state_dict` to match saved param_groups onto +the current optimizer's param_groups during checkpoint resume. It MUST cover +every per-group field that influences scheduler or optimizer behavior; +otherwise two groups distinguishable only by, say, ``max_lr`` collide in the +matching dict and the second one's config silently overwrites the first's +on load — producing wrong LRs after restart and (on a converged-enough model) +loss explosion at the next optimizer step. + +These tests pin the identifier composition and verify the matching tolerates +keys that aren't present on every group (e.g. ``start_wd``/``end_wd``/ +``optimizer`` from :class:`ParamGroupOverride` are only set on groups that +explicitly override them). +""" + +from megatron.core.optimizer.optimizer import MegatronOptimizer, param_group_identifier_keys +from megatron.core.optimizer_param_scheduler import ParamGroupOverride + + +def _make_pg(**kw) -> dict: + """Build a minimal param_group dict for matcher tests. ``params`` is required.""" + pg = {"params": []} + pg.update(kw) + return pg + + +def test_identifier_keys_cover_all_param_group_override_fields(): + """REGRESSION: every field declared on ``ParamGroupOverride`` must be in + the identifier. If someone adds a new field (per-group user-facing config), + it must also appear in ``param_group_identifier_keys`` — otherwise two + groups distinguishable only by that field will collide on load. + + We use ``__annotations__.keys()`` as the source of truth for "declared + fields". For any TypedDict that equals ``__required_keys__ | __optional_keys__``, + so the test is invariant to the TypedDict's ``total=`` setting or to a + future split between ``Required[]`` / ``NotRequired[]`` wrappers — see + the docstring on ``_param_group_override_keys`` in ``optimizer.py``. + """ + declared = set(ParamGroupOverride.__annotations__.keys()) + missing = declared - set(param_group_identifier_keys) + assert not missing, ( + f"ParamGroupOverride fields not in param_group_identifier_keys: {missing}. " + f"Either add to identifier_keys, or argue why this field should NOT participate " + f"in save/load matching." + ) + + +def test_identifier_keys_invariant_to_totality(): + """The identifier-key derivation must work regardless of whether + ``ParamGroupOverride`` is declared ``total=False`` (current), + ``total=True``, or mixed via ``Required[]`` / ``NotRequired[]``. + + This guards against a future maintainer flipping the totality setting and + silently emptying the identifier (which would re-introduce the LR-restart + bug class). + """ + # Verify every field is captured by __annotations__ (the source we use). + assert set(ParamGroupOverride.__annotations__.keys()) >= { + 'max_lr', + 'min_lr', + 'start_wd', + 'end_wd', + 'wd_mult', + 'optimizer', + }, ( + "ParamGroupOverride lost expected fields. If a field was renamed, update " + "this test AND any consumers of the identifier." + ) + # Verify the identifier-derivation function returns exactly the annotation set. + from megatron.core.optimizer.optimizer import _param_group_override_keys + + assert set(_param_group_override_keys()) == set(ParamGroupOverride.__annotations__.keys()), ( + "_param_group_override_keys() must return ParamGroupOverride.__annotations__ verbatim " + "so the identifier survives totality / Required[] / NotRequired[] changes." + ) + + +def test_identifier_excludes_mutable_per_step_keys(): + """REGRESSION: keys the scheduler / inner optimizer rewrite every step must NOT + appear in the identifier. If they did, the freshly-built optimizer (at step 0) + would have different values than the saved optimizer (at step N) for the same + logical group, and matching across save/load would fail. + + Specifically: + - ``lr`` is rewritten by ``OptimizerParamScheduler.step`` every iter. + - ``weight_decay`` is rewritten by ``OptimizerParamScheduler.step`` every iter. + - ``step`` is incremented by the inner Adam each call. + + None of these are on ``ParamGroupOverride`` today, so the current + ``param_group_identifier_keys`` derivation is safe. This test pins the assumption + so a future maintainer who (e.g.) adds ``lr`` to ``ParamGroupOverride`` for some + new use case will see the test fail immediately rather than silently regressing + save/load matching. + """ + mutating = {"lr", "weight_decay", "step"} + overlap = mutating.intersection(param_group_identifier_keys) + assert not overlap, ( + f"param_group_identifier_keys contains keys that mutate every optimizer step: " + f"{overlap}. These will differ between save (iter N) and the freshly-built " + f"optimizer (iter 0), causing the matcher to fail to pair saved groups with " + f"current groups. Either remove them from the identifier or, if they're " + f"declared on ParamGroupOverride, mask them out in _param_group_override_keys()." + ) + + +def test_identifier_keys_include_structural_flags(): + """``lr_mult`` / ``is_expert_parallel`` / ``is_decoupled_lr`` are set on every + param_group at construction time and must remain in the identifier so saves + from one process layout match loads in another. + """ + for key in ("lr_mult", "is_expert_parallel", "is_decoupled_lr"): + assert key in param_group_identifier_keys, ( + f"{key!r} missing from param_group_identifier_keys; this would let groups " + f"collide across e.g. EP-on/EP-off boundaries on resume." + ) + + +def test_filter_reorder_distinguishes_groups_by_max_lr(): + """REGRESSION: two groups that differ ONLY by max_lr/min_lr must be matched + correctly across save/load. Pre-fix, both would have produced the same legacy + 4-tuple ``(wd_mult, lr_mult, is_expert_parallel, is_decoupled_lr)`` of + ``(1.0, 1.0, False, False)`` and collided in the matching dict — the second + saved group's config (max_lr) silently clobbered the first's at load time, + producing wrong LRs at the next optimizer step. + """ + # Two current groups with same wd/structural flags but different max_lr — + # this is exactly the recipe pattern (trunk-WD vs. projector-WD). + current = [ + _make_pg( + wd_mult=1.0, + lr_mult=1.0, + is_expert_parallel=False, + is_decoupled_lr=False, + max_lr=2e-5, + min_lr=2e-6, + ), + _make_pg( + wd_mult=1.0, + lr_mult=1.0, + is_expert_parallel=False, + is_decoupled_lr=False, + max_lr=5e-4, + min_lr=5e-5, + ), + ] + # Saved groups (deliberately reordered to exercise the reorder logic). + saved = [ + _make_pg( + wd_mult=1.0, + lr_mult=1.0, + is_expert_parallel=False, + is_decoupled_lr=False, + max_lr=5e-4, + min_lr=5e-5, + # Some recognizable extra field to confirm the right saved group was matched. + _tag="from_projector", + ), + _make_pg( + wd_mult=1.0, + lr_mult=1.0, + is_expert_parallel=False, + is_decoupled_lr=False, + max_lr=2e-5, + min_lr=2e-6, + _tag="from_trunk", + ), + ] + + reordered = MegatronOptimizer._filter_and_reorder_param_groups(current, saved) + + assert len(reordered) == 2 + # current[0] has max_lr=2e-5 → must match the saved group with max_lr=2e-5. + assert reordered[0]["max_lr"] == 2e-5 + assert reordered[0]["_tag"] == "from_trunk" + # current[1] has max_lr=5e-4 → must match the saved group with max_lr=5e-4. + assert reordered[1]["max_lr"] == 5e-4 + assert reordered[1]["_tag"] == "from_projector" + + +def test_filter_reorder_tolerates_missing_optional_keys(): + """Some identifier keys (``start_wd`` / ``end_wd`` / ``optimizer``) come from + ``ParamGroupOverride`` and are only present on groups that explicitly + override them. Default groups don't carry these keys at all, so the matcher + must tolerate missing keys (rather than KeyError-ing). Two groups + missing the same set of keys must remain matchable. + """ + # Both groups have only the always-present keys; they should match by tuple + # of values plus the same None placeholder for missing keys. + current = [ + _make_pg( + wd_mult=1.0, + lr_mult=1.0, + is_expert_parallel=False, + is_decoupled_lr=False, + max_lr=1e-3, + min_lr=1e-4, + # NOT setting start_wd, end_wd, optimizer — these are absent. + ) + ] + saved = [ + _make_pg( + wd_mult=1.0, + lr_mult=1.0, + is_expert_parallel=False, + is_decoupled_lr=False, + max_lr=1e-3, + min_lr=1e-4, + _tag="saved_match", + ) + ] + reordered = MegatronOptimizer._filter_and_reorder_param_groups(current, saved) + assert len(reordered) == 1 + assert reordered[0]["_tag"] == "saved_match" + + +def test_filter_reorder_treats_missing_and_none_identifier_values_as_same(): + """A missing identifier key and an explicit ``None`` value use the same convention.""" + common = dict( + wd_mult=1.0, lr_mult=1.0, is_expert_parallel=False, is_decoupled_lr=False, max_lr=1e-3 + ) + current = [_make_pg(**common, min_lr=None)] + saved = [_make_pg(**common, _tag="saved_match")] + + reordered = MegatronOptimizer._filter_and_reorder_param_groups(current, saved) + + assert reordered[0]["_tag"] == "saved_match" + + +def test_filter_reorder_distinguishes_by_optional_override_key(): + """When a group sets a key that another doesn't (e.g. ``start_wd``), the + identifier tuple must reflect that — the two groups must be distinguishable + rather than collapsed into one match. + """ + common = dict( + wd_mult=1.0, + lr_mult=1.0, + is_expert_parallel=False, + is_decoupled_lr=False, + max_lr=1e-3, + min_lr=1e-4, + ) + current = [ + _make_pg(**common, start_wd=0.05), # explicit per-group start_wd + _make_pg(**common), # default start_wd (absent -> None) + ] + saved = [ + _make_pg(**common, _tag="default_wd"), + _make_pg(**common, start_wd=0.05, _tag="explicit_wd"), + ] + reordered = MegatronOptimizer._filter_and_reorder_param_groups(current, saved) + assert reordered[0]["_tag"] == "explicit_wd" + assert reordered[0].get("start_wd") == 0.05 + assert reordered[1]["_tag"] == "default_wd" + assert "start_wd" not in reordered[1] + + +def test_filter_reorder_handles_nemo_pre_prefix(): + """NeMo renames ``lr_mult``/``wd_mult`` to ``pre_lr_mult``/``pre_wd_mult``. + The matcher's per-key fallback must look up ``pre_`` if ```` is + missing — this preserves NeMo-saved checkpoint compatibility. + """ + common = dict(is_expert_parallel=False, is_decoupled_lr=False, max_lr=1e-3, min_lr=1e-4) + # Current uses standard names, saved uses NeMo's pre_-prefixed names. + current = [_make_pg(**common, wd_mult=1.0, lr_mult=1.0)] + saved = [_make_pg(**common, pre_wd_mult=1.0, pre_lr_mult=1.0, _tag="from_nemo")] + reordered = MegatronOptimizer._filter_and_reorder_param_groups(current, saved) + assert reordered[0]["_tag"] == "from_nemo" diff --git a/tests/unit_tests/pipeline_parallel/test_pipeline_layout.py b/tests/unit_tests/pipeline_parallel/test_pipeline_layout.py index 1c998181b50..7ded4abd1a5 100644 --- a/tests/unit_tests/pipeline_parallel/test_pipeline_layout.py +++ b/tests/unit_tests/pipeline_parallel/test_pipeline_layout.py @@ -140,6 +140,7 @@ def create_args(): args.vocab_file = None args.add_position_embedding = False args.ckpt_assume_constant_structure = False + args.stream_ckpt_dequant = True args.ckpt_load_validate_sharding_integrity = True args.dist_ckpt_strictness = "assume_ok_unexpected" args.fp16 = False diff --git a/tests/unit_tests/resharding/test_copy_services.py b/tests/unit_tests/resharding/test_copy_services.py index 0fe9a40bf60..9e255f9fd9a 100644 --- a/tests/unit_tests/resharding/test_copy_services.py +++ b/tests/unit_tests/resharding/test_copy_services.py @@ -16,6 +16,7 @@ SendOp, match_local_ops_by_task_id, ) +from megatron.core.resharding.copy_services.nixl_copy_service import NixlCopyService def _t(): @@ -164,3 +165,8 @@ def close(self): assert svc.closed is False svc.close() assert svc.closed is True + + +def test_nixl_service_skips_redundant_process_group_barrier(): + """NIXL's ready/data protocol provides its own peer completion.""" + assert NixlCopyService.requires_process_group_barrier is False diff --git a/tests/unit_tests/resharding/test_execution.py b/tests/unit_tests/resharding/test_execution.py index 70d5d14dc03..5323efab82c 100644 --- a/tests/unit_tests/resharding/test_execution.py +++ b/tests/unit_tests/resharding/test_execution.py @@ -6,7 +6,7 @@ non-collocated mode handling. Requires CUDA (uses torch.cuda.synchronize). """ -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch import pytest import torch @@ -262,6 +262,17 @@ def test_transform_not_called_for_non_matching(self): class TestEdgeCases: """Test edge cases in execute_reshard_plan.""" + def test_service_can_skip_process_group_barrier(self): + """A self-synchronizing backend does not use the executor's barrier.""" + Utils.initialize_distributed() + service = MockCopyService() + service.requires_process_group_barrier = False + + with patch("megatron.core.resharding.execution.dist.barrier") as barrier: + execute_reshard_plan(ReshardPlan(send_ops=[], recv_ops=[]), None, None, service) + + barrier.assert_not_called() + def test_empty_plan(self): """Empty plan (no ops) should complete without error.""" plan = ReshardPlan(send_ops=[], recv_ops=[]) diff --git a/tests/unit_tests/resharding/test_planner.py b/tests/unit_tests/resharding/test_planner.py index 55c4019dc09..fb8fcd6e9f9 100644 --- a/tests/unit_tests/resharding/test_planner.py +++ b/tests/unit_tests/resharding/test_planner.py @@ -11,10 +11,13 @@ import pytest +import megatron.core.resharding.planner as planner from megatron.core.resharding.planner import ( _build_descriptors_for_param, _finalize_dp_transfers, _plan_tp, + build_plan_from_rosters, + index_metadata_rosters, ) from megatron.core.resharding.utils import ParameterMetadata, ShardingDescriptor @@ -360,3 +363,100 @@ def test_missing_tp_ranks(self): descs = _build_descriptors_for_param(src, dst) assert descs == [] + + +# =========================================================================== +# build_plan_from_rosters (local, deterministic planning + node-add stability) +# =========================================================================== + + +def _plan_edges(plans): + """Collect (task_id, src_rank, dst_rank) transfers from a {rank: ReshardPlan}. + + Reads them once from every send op and once from every recv op; the two sets + must be equal for the plan to be consistent (a matching, same-task_id recv for + every send). + """ + sends = {(op.task_id, r, op.peer_rank) for r, p in plans.items() for op in p.send_ops} + recvs = {(op.task_id, op.peer_rank, r) for r, p in plans.items() for op in p.recv_ops} + return sends, recvs + + +def _build_all(gathered_pairs): + """Build every rank's plan from a rank-ordered list of (src_meta, dst_meta).""" + dst_by_rank, src_by_name = index_metadata_rosters(gathered_pairs) + return {rank: build_plan_from_rosters(dst_by_rank, src_by_name, rank) for rank in dst_by_rank} + + +def _recv_sig(plan): + """Identity of a plan's recv ops (task_id + slices), for stability comparisons.""" + return [(op.task_id, op.peer_rank, op.my_slice, op.peer_slice) for op in plan.recv_ops] + + +class TestBuildPlanFromRosters: + """Local plan building replayed independently per rank.""" + + def test_task_ids_match_across_ranks(self): + """Sender and receiver, planned independently, agree on task_id per transfer. + + rank 0 sources a replicated weight; ranks 1 and 2 each receive a full copy. + """ + gathered = [ + ([_meta(owner_rank=0, tp_ranks=[0], dp_ranks=[0])], []), # rank 0: source + ([], [_meta(owner_rank=1, tp_ranks=[1], dp_ranks=[1])]), # rank 1: dest + ([], [_meta(owner_rank=2, tp_ranks=[2], dp_ranks=[2])]), # rank 2: dest + ] + plans = _build_all(gathered) + sends, recvs = _plan_edges(plans) + + # Every send has a matching recv with the same task_id, and vice versa. + assert sends == recvs + # Two transfers: 0->1 and 0->2, with distinct task_ids. + assert len(sends) == 2 + assert {(s, d) for _, s, d in sends} == {(0, 1), (0, 2)} + assert len({tid for tid, _, _ in sends}) == 2 + + def test_node_add_keeps_existing_task_ids_stable(self): + """Appending a rank rebuilds locally without renumbering existing transfers.""" + base = [ + ([_meta(owner_rank=0, tp_ranks=[0], dp_ranks=[0])], []), + ([], [_meta(owner_rank=1, tp_ranks=[1], dp_ranks=[1])]), + ([], [_meta(owner_rank=2, tp_ranks=[2], dp_ranks=[2])]), + ] + before = _build_all(base) + + # A new destination rank 3 joins; everyone replays over the grown roster. + grown = base + [([], [_meta(owner_rank=3, tp_ranks=[3], dp_ranks=[3])])] + after = _build_all(grown) + + sends_after, recvs_after = _plan_edges(after) + assert sends_after == recvs_after + # Existing receivers keep the exact same recv ops (task_id + slices). + for rank in (1, 2): + assert _recv_sig(before[rank]) == _recv_sig(after[rank]) + # The new rank added exactly one transfer with a fresh task_id. + assert len(sends_after) == 3 + assert {(s, d) for _, s, d in sends_after} == {(0, 1), (0, 2), (0, 3)} + + +def test_centralized_planner_compatibility_wrapper(monkeypatch): + """The previous public planner name warns and forwards every argument.""" + sentinel = object() + forwarded = {} + + def fake_local(src_module, dst_module, **kwargs): + forwarded["args"] = (src_module, dst_module) + forwarded["kwargs"] = kwargs + return sentinel + + monkeypatch.setattr(planner, "build_local_reshard_plan", fake_local) + with pytest.warns(DeprecationWarning, match="build_local_reshard_plan"): + result = planner.build_centralized_reshard_plan( + "src", "dst", num_experts=8, group="group", src_rank_offset=3, dst_rank_offset=7 + ) + + assert result is sentinel + assert forwarded == { + "args": ("src", "dst"), + "kwargs": {"num_experts": 8, "group": "group", "src_rank_offset": 3, "dst_rank_offset": 7}, + } diff --git a/tests/unit_tests/rl/test_rl_utils.py b/tests/unit_tests/rl/test_rl_utils.py index dd6c85b2125..c37d5ec00e1 100644 --- a/tests/unit_tests/rl/test_rl_utils.py +++ b/tests/unit_tests/rl/test_rl_utils.py @@ -14,7 +14,10 @@ from megatron.core.models.common.language_module.language_module import LanguageModule from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec from megatron.core.models.gpt.gpt_model import GPTModel -from megatron.core.num_microbatches_calculator import destroy_num_microbatches_calculator +from megatron.core.num_microbatches_calculator import ( + destroy_num_microbatches_calculator, + get_num_microbatches, +) from megatron.core.optimizer import OptimizerConfig, get_megatron_optimizer from megatron.core.pipeline_parallel import get_forward_backward_func from megatron.core.pipeline_parallel.utils import is_pp_first_stage, is_pp_last_stage @@ -89,6 +92,22 @@ def detokenize(self, tokens): return [str(tok) for tok in tokens] +def make_token_rollout(trajectory, logprobs, generation_mask=None, reward=1.0, problem_id="p"): + """TokenRollout with the per-turn staleness boilerplate derived from the turn count.""" + turns = len(trajectory) + return TokenRollout( + trajectory=trajectory, + reward=reward, + generation_mask=generation_mask, + logprobs=logprobs, + env_id='MEGAENV', + problem_id=problem_id, + policy_epoch=[[(0, 0)]] * turns, + kv_cache_epoch=[[(0, 0)]] * turns, + num_evictions=[0] * turns, + ) + + class DummyLangModule: def __init__(self, config): self.config = config @@ -390,118 +409,89 @@ def test_get_logprobs(self, initialize_model_parallel, use_sequence_packing): else: assert logprobs.shape == (BATCH, SEQ, VOCAB) - def test_grpo_loss_calculation_all_pi_eq(self): - # All policies are equal: clamping is inactive, ratios are ones. - current_logprobs = torch.ones(BATCH, SEQ) - old_logprobs = torch.ones(BATCH, SEQ) - ref_logprobs = torch.ones(BATCH, SEQ) - advantages = torch.zeros(BATCH) - loss, kl_term, ratios, entropy_term, _, _ = rl_utils.calculate_grpo_loss( - current_logprobs=current_logprobs, - old_logprobs=old_logprobs, - ref_logprobs=ref_logprobs, - advantages=advantages, - clamp_eps_lower=0.1, - clamp_eps_upper=0.1, - kl_beta=0.1, - entropy_weight=0.0, - ) - torch.testing.assert_close(loss, torch.zeros_like(loss)) - torch.testing.assert_close(kl_term, torch.zeros_like(kl_term)) - torch.testing.assert_close(ratios, torch.ones_like(ratios)) - torch.testing.assert_close(entropy_term, -torch.ones_like(ratios) * torch.e) - - def test_grpo_loss_calculation_2x_ratios(self): - # All policies are equal: clamping is inactive, ratios are ones. - current_logprobs = torch.ones(BATCH, SEQ) - old_logprobs = torch.ones(BATCH, SEQ) - torch.log(torch.tensor([2.0])) - ref_logprobs = torch.ones(BATCH, SEQ) - advantages = torch.ones(BATCH) - loss, kl_term, ratios, _, _, _ = rl_utils.calculate_grpo_loss( - current_logprobs=current_logprobs, - old_logprobs=old_logprobs, - ref_logprobs=ref_logprobs, - advantages=advantages, - clamp_eps_lower=2.1, - clamp_eps_upper=2.1, - kl_beta=0.0, - entropy_weight=0.0, - ) - # Clamping does not affect us, as 2.1 [eps] > 2 [ratio]. - # kl_beta = 0 -> we only have the non-kl term of the loss active. - torch.testing.assert_close(loss, -torch.ones_like(loss) * 2) - # pi and pi_{ref} are the same here. - torch.testing.assert_close(kl_term, torch.zeros_like(kl_term)) - # Current probs are 2x more probable than old pi. - torch.testing.assert_close(ratios, torch.ones_like(ratios) * 2) - - def test_entropy_calculation(self): - # All policies are equal: clamping is inactive, ratios are ones. - current_logprobs = torch.ones(BATCH, SEQ) - old_logprobs = torch.ones(BATCH, SEQ) - ref_logprobs = torch.ones(BATCH, SEQ) - advantages = torch.zeros(BATCH) - loss, _, ratios, entropy_term, _, _ = rl_utils.calculate_grpo_loss( - current_logprobs=current_logprobs, - old_logprobs=old_logprobs, - ref_logprobs=ref_logprobs, - advantages=advantages, - clamp_eps_lower=0.1, - clamp_eps_upper=0.1, - kl_beta=0.0, - entropy_weight=1.0, - ) - torch.testing.assert_close(loss, torch.ones_like(ratios) * torch.e) - torch.testing.assert_close(entropy_term, -torch.ones_like(ratios) * torch.e) - - def test_grpo_loss_truncation(self): - # All ratios are 2 - _, _, _, _, truncated_from_above, truncated_from_below = rl_utils.calculate_grpo_loss( + @pytest.mark.parametrize( + "ratio, advantage, clamp_eps, kl_beta, entropy_weight, expected", + [ + # All policies equal: clamping inactive, unit ratios, zero loss and kl. + pytest.param( + 1.0, + 0.0, + 0.1, + 0.1, + 0.0, + dict(loss=0.0, kl_term=0.0, ratios=1.0, entropy_term=-torch.e), + id="all_pi_eq", + ), + # Current probs 2x old; eps 2.1 > ratio keeps clamping inactive; kl_beta 0 leaves + # only the policy term, and pi == pi_ref keeps the kl term zero anyway. + pytest.param( + 2.0, 1.0, 2.1, 0.0, 0.0, dict(loss=-2.0, kl_term=0.0, ratios=2.0), id="2x_ratios" + ), + # kl_beta 0, entropy_weight 1: the loss is exactly the negated entropy term. + pytest.param( + 1.0, 0.0, 0.1, 0.0, 1.0, dict(loss=torch.e, entropy_term=-torch.e), id="entropy" + ), + ], + ) + def test_grpo_loss_calculation( + self, ratio, advantage, clamp_eps, kl_beta, entropy_weight, expected + ): + outputs = rl_utils.calculate_grpo_loss( current_logprobs=torch.ones(BATCH, SEQ), - old_logprobs=0.5 * torch.ones(BATCH, SEQ), + old_logprobs=torch.ones(BATCH, SEQ) - torch.log(torch.tensor(ratio)), ref_logprobs=torch.ones(BATCH, SEQ), - advantages=torch.zeros(BATCH), - clamp_eps_lower=0.1, - clamp_eps_upper=0.1, - kl_beta=0.1, - entropy_weight=0.0, + advantages=torch.full((BATCH,), advantage), + clamp_eps_lower=clamp_eps, + clamp_eps_upper=clamp_eps, + kl_beta=kl_beta, + entropy_weight=entropy_weight, ) - assert truncated_from_above.float().mean() == 1 - assert truncated_from_below.float().sum() == 0 + outputs = dict(zip(("loss", "kl_term", "ratios", "entropy_term"), outputs)) + for name, want in expected.items(): + torch.testing.assert_close(outputs[name], torch.full_like(outputs[name], want)) - # All ratios are 0.01 - _, _, _, _, truncated_from_above, truncated_from_below = rl_utils.calculate_grpo_loss( - current_logprobs=0.01 * torch.ones(BATCH, SEQ), - old_logprobs=torch.ones(BATCH, SEQ), - ref_logprobs=torch.ones(BATCH, SEQ), - advantages=torch.zeros(BATCH), - clamp_eps_lower=0.1, - clamp_eps_upper=0.1, - kl_beta=0.1, - entropy_weight=0.0, - ) - assert truncated_from_above.float().sum() == 0 - assert truncated_from_below.float().mean() == 1 - - # Mixed ratios: [[2., 0.5], [20., 1.]] - current_logprobs = torch.tensor([[1.0, 1.0], [1.0, 1.0]]) - old_logprobs = torch.tensor([[0.5, 2.0], [0.05, 1.0]]) - _, _, _, _, truncated_from_above, truncated_from_below = rl_utils.calculate_grpo_loss( - current_logprobs=current_logprobs, - old_logprobs=old_logprobs, - ref_logprobs=old_logprobs, - advantages=torch.zeros(BATCH), + @pytest.mark.parametrize( + "current, old, expected_above, expected_below", + [ + # Ratios uniformly above the clamp window: everything truncates from above. + pytest.param( + torch.ones(BATCH, SEQ), + 0.5 * torch.ones(BATCH, SEQ), + torch.full((BATCH, SEQ), True), + torch.full((BATCH, SEQ), False), + id="all_above", + ), + # Ratios uniformly below the clamp window: everything truncates from below. + pytest.param( + 0.01 * torch.ones(BATCH, SEQ), + torch.ones(BATCH, SEQ), + torch.full((BATCH, SEQ), False), + torch.full((BATCH, SEQ), True), + id="all_below", + ), + # Mixed: above, below, above, and a unit ratio truncates neither way. + pytest.param( + torch.ones(2, 2), + torch.tensor([[0.5, 2.0], [0.05, 1.0]]), + torch.tensor([[True, False], [True, False]]), + torch.tensor([[False, True], [False, False]]), + id="mixed", + ), + ], + ) + def test_grpo_loss_truncation(self, current, old, expected_above, expected_below): + *_, truncated_from_above, truncated_from_below = rl_utils.calculate_grpo_loss( + current_logprobs=current, + old_logprobs=old, + ref_logprobs=old, + advantages=torch.zeros(current.shape[0]), clamp_eps_lower=0.1, clamp_eps_upper=0.1, kl_beta=0.1, entropy_weight=0.0, ) - torch.testing.assert_close( - truncated_from_above, torch.tensor([[True, False], [True, False]]) - ) - torch.testing.assert_close( - truncated_from_below, torch.tensor([[False, True], [False, False]]) - ) + torch.testing.assert_close(truncated_from_above, expected_above) + torch.testing.assert_close(truncated_from_below, expected_below) @pytest.mark.parametrize( "initialize_model_parallel", @@ -513,7 +503,8 @@ def test_grpo_loss_truncation(self): indirect=["initialize_model_parallel"], ) def test_prepare_data_for_update(self, initialize_model_parallel): - """Test that getting logprobs at least does not crash.""" + """Logprobs path runs; single-turn EOD guard holds; multi-turn turns are split, + padded to a DP*microbatch multiple, and size the microbatch calculator by turns.""" world_size, dp, tp, pp = initialize_model_parallel # Here I assume that we will be consuming all data in one step. group_size = 2 @@ -531,59 +522,63 @@ def test_prepare_data_for_update(self, initialize_model_parallel): model = MockModel() tokenizer = MockTokenizer() - r1 = TokenRollout( - trajectory=[[1, 2, 3]], - reward=3.14, - generation_mask=[[False, True, True]], - logprobs=[[0.1, 0.2, 0.3]], - env_id='MEGAENV', - problem_id="2", - policy_epoch=[[(0, 0)]], - kv_cache_epoch=[[(0, 0)]], - num_evictions=[0], + # A single-turn rollout whose only turn is short and lacks eod must be rejected: + # a single-turn completion has no tool-call boundary to justify stopping early. + bad = make_token_rollout( + [[1, 2, 3]], [[0.1, 0.2, 0.3]], [[False, True, True]], reward=3.14, problem_id="2" ) - r2 = TokenRollout( - trajectory=[[1, 2, 3, 4]], - reward=0.14, - generation_mask=[[False, True, True, True]], - logprobs=[[0.1, 0.2, 0.3, -1.2]], - env_id='MEGAENV', - problem_id="2", - policy_epoch=[[(0, 0)]], - kv_cache_epoch=[[(0, 0)]], - num_evictions=[0], - ) - - rollouts = [[r1, r2] for _ in range(dp)] - try: + with pytest.raises(AssertionError, match="must end in eod"): rl_utils.prepare_data_for_update( - [model], {}, rollouts, tokenizer, sequence_packing=False, is_correction=False + [model], + {}, + [[bad] for _ in range(dp)], + tokenizer, + sequence_packing=False, + is_correction=False, ) - except AssertionError as e: - # We expect trajectories to come padded there. - assert str(e).startswith('Rollout is not the correct length') - r1 = TokenRollout( - trajectory=torch.tensor([[1, 2, 3, tokenizer.eod]], dtype=torch.float).cuda(), + # Multi-turn rollouts with uneven turn counts: a turn may stop on a tool-call + # boundary (short, no eod) and is accepted. 2 + 3 turns per group * dp groups + # = 5*dp turns, padded up to the next multiple of micro_batch_size*dp (= 2*dp) + # -> 6*dp turns. With samples_ratio 1 the calculator is sized by total turns + # (6*dp -> 3 microbatches), not by rollout count (the pre-fix bug gave 1). + mt1 = make_token_rollout( + [[1, 2, 3], [1, 2, 3, 4]], + [[-0.1, -0.2], [-0.3, -0.4]], + [[False, True, True], [False, False, True, True]], + problem_id="1", + ) + mt2 = make_token_rollout( + [[1, 2], [1, 2, 3], [1, 2, 3, 4]], + [[-0.1], [-0.2], [-0.3]], + [[False, True], [False, False, True], [False, False, False, True]], + reward=0.0, + problem_id="3", + ) + rl_utils.prepare_data_for_update( + [model], + {}, + [[mt1, mt2] for _ in range(dp)], + tokenizer, + sequence_packing=False, + is_correction=False, + ) + # 5*dp turns padded to 6*dp; 6*dp / (micro_batch_size 2 * dp) = 3 microbatches. + assert get_num_microbatches() == 3 + + r1 = make_token_rollout( + torch.tensor([[1, 2, 3, tokenizer.eod]], dtype=torch.float).cuda(), + torch.tensor([[-0.2, -0.3, -3.2]]).cuda(), + torch.tensor([[False, True, True, True]], dtype=torch.float).cuda(), reward=3.14, - generation_mask=torch.tensor([[False, True, True, True]], dtype=torch.float).cuda(), - logprobs=torch.tensor([[-0.2, -0.3, -3.2]]).cuda(), - env_id='MEGAENV', problem_id="2", - policy_epoch=[[(0, 0)]], - kv_cache_epoch=[[(0, 0)]], - num_evictions=[0], ) - r2 = TokenRollout( - trajectory=torch.tensor([[1, 2, 234, tokenizer.eod]], dtype=torch.float).cuda(), + r2 = make_token_rollout( + torch.tensor([[1, 2, 234, tokenizer.eod]], dtype=torch.float).cuda(), + torch.tensor([[-0.2, -0.3, -1.2]]), + torch.tensor([[False, True, True, True]], dtype=torch.float).cuda(), reward=0.14, - generation_mask=torch.tensor([[False, True, True, True]], dtype=torch.float).cuda(), - logprobs=torch.tensor([[-0.2, -0.3, -1.2]]), - env_id='MEGAENV', problem_id="2", - policy_epoch=[[(0, 0)]], - kv_cache_epoch=[[(0, 0)]], - num_evictions=[0], ) rollouts = [[r1, r2] for _ in range(dp)] data_iter, _, _ = rl_utils.prepare_data_for_update( @@ -595,74 +590,96 @@ def test_prepare_data_for_update(self, initialize_model_parallel): # All probabilities should be uniform. torch.testing.assert_close(old_logprobs.exp(), torch.ones_like(old_logprobs) / VOCAB) - @pytest.mark.parametrize("use_sequence_packing", [True, False]) - @pytest.mark.parametrize("num_turns", [1, 2]) - def test_prepare_trajectories(self, use_sequence_packing, num_turns): - """Test that rollouts are properly prepared for training.""" - seq_length = 8 + @pytest.mark.parametrize( + "initialize_model_parallel", + [pytest.param((1, 1), id="tp1-pp1")], + indirect=["initialize_model_parallel"], + ) + def test_prepare_data_for_update_oversampling(self, initialize_model_parallel): + """Oversampling (ratio < 1) consumes a fraction of the (padded) turn count per step: + the microbatch calculator is sized by ceil(ratio * total turns), not the full batch.""" + world_size, dp, tp, pp = initialize_model_parallel + tokenizer = MockTokenizer() + model = MockModel() + + # ratio = global_batch_size/(prompts*group) = 2*dp/(dp*4) = 0.5. + # 4*dp single-turn turns (already a multiple of 2*dp); ceil(0.5 * 4*dp) = 2*dp; + # 2*dp / (2 * dp) = 1 microbatch. self.create_test_args( - rl_use_sequence_packing=use_sequence_packing, - rl_sequence_packing_bin_size=20, - rl_skip_bos_token=False, - micro_batch_size=1, - seq_length=seq_length, + micro_batch_size=2, + seq_length=4, + curr_iteration=1, + tensor_model_parallel_size=tp, + pipeline_model_parallel_size=pp, + global_batch_size=dp * 2, + grpo_prompts_per_step=dp, + grpo_group_size=4, + ) + + def single(problem_id, reward): + return make_token_rollout( + [[1, 2, 3, tokenizer.eod]], + [[-0.1, -0.2, -0.3]], + [[False, True, True, True]], + reward=reward, + problem_id=problem_id, + ) + + rollouts = [[single(str(i), float(i % 2)) for i in range(4)] for _ in range(dp)] + rl_utils.prepare_data_for_update( + [model], {}, rollouts, tokenizer, sequence_packing=False, is_correction=False ) + assert get_num_microbatches() == 1 + + @pytest.mark.parametrize("num_turns", [1, 2]) + def test_prepare_trajectories(self, num_turns): + """Each (rollout, turn_idx) unit becomes one padded training row (a rollout with T + turns contributes T rows); PAD_TURN_UNIT entries become inert all-pad rows with no + generated tokens and no inference logprobs (DP count equalization).""" + seq_length = 8 + self.create_test_args(rl_skip_bos_token=False, micro_batch_size=1, seq_length=seq_length) tokenizer = MockTokenizer() + eod, pad = tokenizer.eod, tokenizer.pad - # Create rollouts of varying lengths - r1 = TokenRollout( - trajectory=[[1, 2, 3, tokenizer.eod]] * num_turns, + r1 = make_token_rollout( + [[1, 2, 3, eod]] * num_turns, + [[0.1, 0.2, 0.3, 0.35]] * num_turns, + [[False, True, True, True]] * num_turns, reward=3.14, - generation_mask=[[False, True, True, True]] * num_turns, - logprobs=[[0.1, 0.2, 0.3, 0.35]] * num_turns, - env_id='MEGAENV', problem_id="1", - policy_epoch=[[(0, 0)]] * num_turns, - kv_cache_epoch=[[(0, 0)]] * num_turns, - num_evictions=[0] * num_turns, ) - r2 = TokenRollout( - trajectory=[[4, 5, 6, 7, tokenizer.eod]] * num_turns, + r2 = make_token_rollout( + [[4, 5, 6, 7, eod]] * num_turns, + [[0.4, 0.5, 0.6, 0.7, 0.75]] * num_turns, + [[False, True, True, True, True]] * num_turns, reward=0.14, - generation_mask=[[False, True, True, True, True]] * num_turns, - logprobs=[[0.4, 0.5, 0.6, 0.7, 0.75]] * num_turns, - env_id='MEGAENV', problem_id="2", - policy_epoch=[[(0, 0)]] * num_turns, - kv_cache_epoch=[[(0, 0)]] * num_turns, - num_evictions=[0] * num_turns, ) - r3 = TokenRollout( - trajectory=[[8, 9, tokenizer.eod]] * num_turns, + r3 = make_token_rollout( + [[8, 9, eod]] * num_turns, + [[0.8, 0.9, 0.95]] * num_turns, + [[False, True, True]] * num_turns, reward=2.71, - generation_mask=[[False, True, True]] * num_turns, - logprobs=[[0.8, 0.9, 0.95]] * num_turns, - env_id='MEGAENV', problem_id="3", - policy_epoch=[[(0, 0)]] * num_turns, - kv_cache_epoch=[[(0, 0)]] * num_turns, - num_evictions=[0] * num_turns, ) - rollouts = [r1, r2, r3] - + turn_units = [ + (rollout, turn_idx) + for rollout in [r1, r2, r3] + for turn_idx in range(len(rollout.trajectory)) + ] + [rl_utils.PAD_TURN_UNIT] trajs, genmask, inference_logprobs = rl_utils.prepare_trajectories( - rollouts, - tokenizer, - seq_length, - sequence_packing=use_sequence_packing, - skip_bos_token=False, + turn_units, tokenizer, seq_length, skip_bos_token=False ) expected_trajs = torch.tensor( - [ - [1, 2, 3, tokenizer.eod] + [tokenizer.pad] * 4, - [4, 5, 6, 7, tokenizer.eod] + [tokenizer.pad] * 3, - [8, 9, tokenizer.eod] + [tokenizer.pad] * 5, - ], + [[1, 2, 3, eod] + [pad] * 4, [4, 5, 6, 7, eod] + [pad] * 3, [8, 9, eod] + [pad] * 5], dtype=torch.long, device=trajs.device, ).repeat_interleave(num_turns, dim=0) + expected_trajs = torch.cat( + [expected_trajs, torch.full((1, seq_length), pad, dtype=expected_trajs.dtype)] + ) assert torch.equal(trajs, expected_trajs) expected_genmask = torch.tensor( @@ -674,75 +691,98 @@ def test_prepare_trajectories(self, use_sequence_packing, num_turns): dtype=torch.bool, device=genmask.device, ).repeat_interleave(num_turns, dim=0) + expected_genmask = torch.cat( + [expected_genmask, torch.zeros((1, seq_length), dtype=torch.bool)] + ) assert torch.equal(genmask, expected_genmask) - if use_sequence_packing: - expected_logprobs = torch.tensor( - [ - [0.1, 0.2, 0.3, 0.35] + [0.0] * 4, - [0.4, 0.5, 0.6, 0.7, 0.75] + [0.0] * 3, - [0.8, 0.9, 0.95] + [0.0] * 5, - ], - dtype=torch.float32, - device=inference_logprobs.device, - ).repeat_interleave(num_turns, dim=0) - torch.testing.assert_close(inference_logprobs, expected_logprobs, rtol=0, atol=0) - else: - expected_logprobs = [ - [0.1, 0.2, 0.3, 0.35], - [0.4, 0.5, 0.6, 0.7, 0.75], - [0.8, 0.9, 0.95], - ] - expected_logprobs = [el for el in expected_logprobs for _ in range(num_turns)] - assert len(inference_logprobs) == len(expected_logprobs) - for got, exp in zip(inference_logprobs, expected_logprobs): - got_t = got if torch.is_tensor(got) else torch.tensor(got, dtype=torch.float32) - exp_t = torch.tensor(exp, dtype=torch.float32, device=got_t.device) - torch.testing.assert_close(got_t, exp_t, rtol=0, atol=0) - - def test_single_turn_advantage_calculation(self): - rewards = [[-1, 1], [4, 4]] - num_turns = [[1, 1], [1, 1]] - advs = rl_utils.calculate_grpo_advantages(rewards, num_turns) - torch.testing.assert_close( - torch.tensor(advs), torch.tensor([-1, 1.0, 0.0, 0.0]), atol=1e-4, rtol=1e-5 - ) + # Per-row list: unpadded tensor per real row, None for the pad unit. (Packing-mode + # densification happens at the call site via _pad_nonnull_with_zeros, tested separately.) + expected_logprobs = [[0.1, 0.2, 0.3, 0.35], [0.4, 0.5, 0.6, 0.7, 0.75], [0.8, 0.9, 0.95]] + expected_logprobs = [el for el in expected_logprobs for _ in range(num_turns)] + [None] + assert len(inference_logprobs) == len(expected_logprobs) + for got, exp in zip(inference_logprobs, expected_logprobs): + if exp is None: + assert got is None + else: + exp_t = torch.tensor(exp, dtype=torch.float32, device=got.device) + torch.testing.assert_close(got, exp_t, rtol=0, atol=0) - def test_multi_turn_advantage_calculation(self): - rewards = [[-1, 1], [4, 4]] - num_turns = [[2, 1], [1, 3]] - advs = rl_utils.calculate_grpo_advantages(rewards, num_turns) - torch.testing.assert_close( - torch.tensor(advs), - torch.tensor([-1, -1, 1.0, 0.0, 0.0, 0.0, 0.0]), - atol=1e-4, - rtol=1e-5, - ) + @pytest.mark.parametrize( + "num_turns, expected", + [ + pytest.param([[1, 1], [1, 1]], [-1.0, 1.0, 0.0, 0.0], id="single_turn"), + # A rollout's group advantage is repeated once per turn. + pytest.param([[2, 1], [1, 3]], [-1.0, -1.0, 1.0, 0.0, 0.0, 0.0, 0.0], id="multi_turn"), + ], + ) + def test_advantage_calculation(self, num_turns, expected): + advs = rl_utils.calculate_grpo_advantages([[-1, 1], [4, 4]], num_turns) + torch.testing.assert_close(torch.tensor(advs), torch.tensor(expected), atol=1e-4, rtol=1e-5) - def test_pad_list_of_nones(self): - with pytest.raises(ValueError) as e_info: - rl_utils._pad_nonnull_with_zeros([None] * 3, 42) - assert "At least one" in str(e_info) + @pytest.mark.parametrize( + "scenario, expected_turn_lens, expected_traj_lens, expected_num_turns", + [ + pytest.param("single_turn_only", [[4, 3]], [[4, 3]], [[1, 1]], id="single_turn_only"), + pytest.param( + "multi_and_single", [[4, 3, 4]], [[7, 4]], [[2, 1]], id="multi_and_single" + ), + ], + ) + def test_compute_group_stats( + self, scenario, expected_turn_lens, expected_traj_lens, expected_num_turns + ): + """Length metrics: single-turn rollouts use the plain per-turn length, while a multi-turn + TokenRollout re-encodes the prior conversation, so its per-turn lengths are reported + incrementally and its trajectory length is the final conversation length (not the inflated + overlap sum).""" + tokenizer = MockTokenizer() + eod = tokenizer.eod - def test_pad_with_wrong_params(self): - with pytest.raises(ValueError) as e_info: - rl_utils._pad_nonnull_with_zeros([torch.zeros(5)], 4) - assert "larger length" in str(e_info) + def single(traj, reward): + return make_token_rollout( + [traj], [[0.0]], [[False] * len(traj)], reward=reward, problem_id="s" + ) - def test_pad_full_size(self): - padded = rl_utils._pad_nonnull_with_zeros([torch.zeros(5), torch.zeros(5)], 5) - assert padded.shape == (2, 5) + if scenario == "single_turn_only": + group = [single([1, 2, 3, eod], 1.0), single([1, 2, eod], 0.0)] + else: + # Cumulative per-turn lengths 4 then 7 -> turn 1 adds 3 tokens; trajectory length is + # the full conversation (7), not 4 + 7 = 11. + multi = make_token_rollout( + [[1, 2, 3, eod], [1, 2, 3, eod, 9, 8, eod]], + [[0.1, 0.2], [0.3, 0.4]], + [[False, False, True, True], [False, False, False, False, False, True, True]], + problem_id="m", + ) + group = [multi, single([1, 2, 3, eod], 0.0)] - def test_pad_some_nones(self): - padded = rl_utils._pad_nonnull_with_zeros([None, torch.zeros(5)], 5) - assert padded.shape == (2, 5) - assert (padded[0] == 0).all() + stats = rl_utils.compute_group_stats([group], tokenizer, seq_len=8) + assert stats.turn_lens == expected_turn_lens + assert stats.traj_lens == expected_traj_lens + assert stats.num_turns == expected_num_turns - def test_pad_normal(self): - padded = rl_utils._pad_nonnull_with_zeros( - [torch.zeros(2), torch.zeros(3), torch.zeros(4)], 5 - ) - assert padded.shape == (3, 5) + @pytest.mark.parametrize( + "lengths, max_len, expected_shape", + [ + pytest.param([2, 3, 4], 5, (3, 5), id="normal"), + pytest.param([5, 5], 5, (2, 5), id="full_size"), + pytest.param([None, 5], 5, (2, 5), id="some_nones"), + # All-None (all-PAD rank): still a zero [num_rows, max_len] tensor, so every DP rank + # produces the same shape and joins the sequence-packing all_gather. + pytest.param([None, None, None], 42, (3, 42), id="all_nones"), + pytest.param([5], 4, "larger length", id="too_long_raises"), + ], + ) + def test_pad_nonnull_with_zeros(self, lengths, max_len, expected_shape): + data = [None if l is None else torch.zeros(l) for l in lengths] + if isinstance(expected_shape, str): + with pytest.raises(ValueError, match=expected_shape): + rl_utils._pad_nonnull_with_zeros(data, max_len) + return + padded = rl_utils._pad_nonnull_with_zeros(data, max_len) + assert padded.shape == expected_shape + assert (padded == 0).all() # zero inputs and zero-filled padding/None rows @pytest.mark.parametrize( "initialize_model_parallel", @@ -1119,7 +1159,10 @@ def test_get_logprobs_cuda_graphs(self, initialize_model_parallel): [pytest.param((1, 1), id="tp1-pp1")], indirect=["initialize_model_parallel"], ) - def test_prep_wandb_metrics(self, initialize_model_parallel): + @pytest.mark.parametrize( + "inject_placeholders", [False, True], ids=["clean", "with_placeholders"] + ) + def test_prep_wandb_metrics(self, initialize_model_parallel, inject_placeholders): # This tests the computation and makes us fail noisily if # inputs assumptions are changed, e.g. we expect rewards to come in groups (list[list[int]]). traj_lens = [[3, 3], [1, 2]] @@ -1134,8 +1177,29 @@ def test_prep_wandb_metrics(self, initialize_model_parallel): completed_epochs = [[5, 3], [5, 1]] num_evictions = [[0, 1], [0, 0]] current_iteration = 6 + if inject_placeholders: + for lst, sentinel in ( + (traj_lens, 0), + (rewards, 0.0), + (num_turns, 0), + (num_evictions, 0), + ): + for group in lst: + group.append(sentinel) + lst.append([sentinel, sentinel]) # the fully failed extra group + for lst in (policy_epoch, kv_cache_epoch): + for group in lst: + group.append([0]) # sentinel epoch stamp of a placeholder + lst.append([[0], [0]]) + # Placeholders contribute no turns, and compute_group_stats already + # excludes them from completed_epochs; the failed group adds empty + # inner lists, which the group-level stats must skip, not crash on. + turn_lens.append([]) + completed_epochs.append([]) + # advantages stay [0, 1]: zero-turn rollouts emit no advantage entries. + writer = MagicMock() metrics = rl_utils.prep_wandb_metrics( - MagicMock(), + writer, traj_lens, turn_lens, rewards, @@ -1147,8 +1211,21 @@ def test_prep_wandb_metrics(self, initialize_model_parallel): num_evictions=num_evictions, current_iteration=current_iteration, ) - assert metrics["mean_reward"] == 0.75 + assert metrics["failed_rollouts/count"] == (4 if inject_placeholders else 0) + assert metrics["failed_rollouts/ratio"] == (0.5 if inject_placeholders else 0.0) + # Reward aggregates keep the placeholder zeros by design: group means + # become [2/3, 1/3, 0] instead of [1, 0.5]. + assert np.isclose(metrics["mean_reward"], 1 / 3 if inject_placeholders else 0.75) assert metrics["mean_advantage"] == 0.5 + # The rollout table lists real rollouts only, in either case. + rollout_table_calls = [ + c + for c in writer.Table.call_args_list + if c.kwargs.get("columns", [None])[:2] == ["reward", "traj_length"] + ] + assert len(rollout_table_calls) == 1 + rows = rollout_table_calls[0].kwargs["data"] + assert [r[3] for r in rows] == [2, 4, 1, 6] # policy_staleness column assert metrics["nonzero_groups_ratio"] == 0.5 assert metrics["max_traj_length"] == 3 assert metrics["min_traj_length"] == 1 @@ -1181,3 +1258,74 @@ def test_prep_wandb_metrics(self, initialize_model_parallel): assert metrics["max_num_evictions"] == 1 # mean_completion_gap = mean([6-5, 6-3, 6-5, 6-1]) = mean([1, 3, 1, 5]) = 2.5 assert metrics["mean_completion_gap"] == 2.5 + + def test_compute_group_stats_excludes_placeholders_from_metric_fields(self): + def real_rollout(tokens, epoch, problem_id): + return TokenRollout( + trajectory=[tokens], + generation_mask=[[True] * len(tokens)], + reward=1.0, + logprobs=[[0.0] * len(tokens)], + env_id="swe", + problem_id=problem_id, + policy_epoch=[[(0, epoch)]], + kv_cache_epoch=[[(0, epoch)]], + num_evictions=[0], + ) + + def placeholder(): + return TokenRollout( + trajectory=[], + generation_mask=[], + reward=0.0, + logprobs=[], + env_id="swe", + problem_id="placeholder", + policy_epoch=[[(0, 0)]], + kv_cache_epoch=[[(0, 0)]], + num_evictions=[0], + ) + + eod = MockTokenizer().eod + rollouts = [ + [ + real_rollout([1, 2, eod], epoch=5, problem_id="p0"), + real_rollout([1, 2, 3, eod], epoch=6, problem_id="p0"), + placeholder(), + ], + [placeholder(), placeholder(), placeholder()], + ] + stats = rl_utils.compute_group_stats(rollouts, MockTokenizer(), seq_len=16) + + # Per-rollout lists keep the placeholder entries: alignment with rewards + # and num_turns is what lets prep_wandb_metrics mask them downstream. + assert stats.num_turns == [[1, 1, 0], [0, 0, 0]] + assert stats.policy_epoch == [[[5], [6], [0]], [[0], [0], [0]]] + assert stats.traj_lens == [[3, 4, 0], [0, 0, 0]] + # Per-turn lists exclude placeholders entirely: no sentinel epoch-0 stamp + # in completed_epochs, no fake 0-length turn for all-placeholder groups. + assert stats.completed_epochs == [[5, 6], []] + assert stats.turn_lens == [[3, 4], []] + # Rewards keep the placeholder zeros (they shape the GRPO baseline). + assert stats.rewards == [[1.0, 1.0, 0.0], [0.0, 0.0, 0.0]] + + # End to end: the sentinel epochs never reach the staleness metrics. + metrics = rl_utils.prep_wandb_metrics( + MagicMock(), + stats.traj_lens, + stats.turn_lens, + stats.rewards, + stats.num_turns, + stats.advantages, + policy_epoch=stats.policy_epoch, + kv_cache_epoch=stats.kv_cache_epoch, + completed_epochs=stats.completed_epochs, + num_evictions=stats.num_evictions, + current_iteration=7, + ) + assert metrics["max_policy_staleness"] == 2 # 7 - 5, not 7 - 0 + assert metrics["min_traj_length"] == 3 + assert metrics["min_num_turns"] == 1 + assert metrics["mean_completion_gap"] == np.mean([2, 1]) + assert metrics["failed_rollouts/count"] == 4 + assert np.isclose(metrics["failed_rollouts/ratio"], 4 / 6) diff --git a/tests/unit_tests/rl/test_grouped_rollouts.py b/tests/unit_tests/rl/test_rollout_generation.py similarity index 50% rename from tests/unit_tests/rl/test_grouped_rollouts.py rename to tests/unit_tests/rl/test_rollout_generation.py index ef80319ea74..b9c26c7600e 100644 --- a/tests/unit_tests/rl/test_grouped_rollouts.py +++ b/tests/unit_tests/rl/test_rollout_generation.py @@ -1,23 +1,26 @@ -# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. import asyncio from unittest.mock import MagicMock import numpy as np import pytest -from pydantic import ValidationError +from pydantic import Field, ValidationError from megatron.rl.agent.api import ( + EpisodeResult, GroupedRolloutGenerator, GroupedRolloutRequest, GroupRolloutParams, Rollout, RolloutGenerator, RolloutRequest, + TokenRollout, + _SubmissionGate, ) from megatron.rl.agent.reward_only_agent import RewardOnlyAgent from megatron.rl.agent.weighted_multi_task import AgentConfig, WeightedMultiTask -from megatron.rl.inference import InferenceResponse, LLMChatMessage, ReturnsRaw +from megatron.rl.inference import InferenceResponse, LLMChatMessage, ReturnsRaw, ReturnsTokens class MockInferenceInterface(ReturnsRaw): @@ -68,22 +71,30 @@ async def prepare_group_rollout(self, request): idx = self._call_count self._call_count += 1 self.prepare_group_rollout_calls += 1 - inference_request = request.inference_interface.prepare_request( - f"t{idx}", request.generation_args - ) - async def build_rollout(response): - response_idx = int(response.response.content.removeprefix("t")) + async def run_episode(): + # Single-turn agent: the episode is one inference on the group's prompt. + turn_request = request.inference_interface.prepare_request( + f"t{idx}", request.generation_args + ) + response = await self.get_rollout_response(request, turn_request) + return EpisodeResult( + responses=[response], conversation=[*turn_request.prompt, response.response] + ) + + async def build_rollout(episode): + responses = episode.responses + reward = float(responses[-1].response.content.removeprefix("t")) return Rollout( - trajectory=[response.raw_text], - reward=float(response_idx), + trajectory=[r.raw_text for r in responses], + reward=reward, env_id=self.env_id, - policy_epoch=[response.policy_epoch], - kv_cache_epoch=[response.kv_cache_epoch], - num_evictions=[response.num_evictions], + policy_epoch=[r.policy_epoch for r in responses], + kv_cache_epoch=[r.kv_cache_epoch for r in responses], + num_evictions=[r.num_evictions for r in responses], ) - return GroupRolloutParams(inference_request=inference_request, build_rollout=build_rollout) + return GroupRolloutParams(run_episode=run_episode, build_rollout=build_rollout) class CountingRewardAgent(RewardOnlyAgent): @@ -103,6 +114,101 @@ async def get_reward(self, response, golden, finish_reason): return float(int(response.removeprefix("t")) == golden["idx"]) +async def _flush(rounds: int = 50): + """Let pipeline stage tasks settle (mock inference is zero-delay).""" + for _ in range(rounds): + await asyncio.sleep(0) + + +class TestSubmissionGate: + @pytest.mark.asyncio + @pytest.mark.parametrize("submission", ["R", "G", "B"]) + async def test_release_requires_matching_granularity(self, submission): + gate = _SubmissionGate(capacity=1, submission=submission) + await gate.acquire_for(submission) + assert gate.held == 1 + for granularity in ("R", "G", "B"): + if granularity == submission: + continue + gate.release_for(granularity) + assert gate.held == 1 + assert gate.release_calls == 0 + gate.release_for(submission) + assert gate.held == 0 + assert gate.release_calls == 1 + + +class TestConsumptionRelease: + """G-submission gate slots must recycle on trainer consumption, not assembly.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "consumption_granularity, num_groups", + [ + pytest.param("G", 1, id="group_consumption"), + pytest.param("B", 2, id="batch_consumption"), + ], + ) + async def test_group_submission_stalls_until_consumption( + self, consumption_granularity, num_groups + ): + capacity = 4 + gen = MockGenerator(parallel_generation_tasks=capacity) + request = GroupedRolloutRequest( + num_groups=num_groups, + rollouts_per_group=1, + inference_interface=MockInferenceInterface(), + streaming=True, + submission_granularity="G", + consumption_granularity=consumption_granularity, + ) + it = gen.get_grouped_rollouts(request) + try: + for pulled in range(1, capacity + 3): + # wait_for turns the deadlock failure mode (a slot never freed) + # into a test failure instead of a hang. + await asyncio.wait_for(anext(it), timeout=10) + await _flush() + # Each yield frees exactly one group slot on the consumer's next + # resume, so submission tracks consumption with a one-slot skew + # (the release for the latest pull hasn't fired yet). On + # assembly-release semantics this runs away unbounded; if no + # consume-site release existed, the loop would deadlock at + # `pulled == capacity + 1`. + assert gen.prepare_group_rollout_calls == capacity + pulled - 1 + finally: + await it.aclose() + + @pytest.mark.asyncio + async def test_batch_submission_releases_once_per_batch(self): + gen = MockGenerator(parallel_generation_tasks=1) + request = GroupedRolloutRequest( + num_groups=2, + rollouts_per_group=1, + inference_interface=MockInferenceInterface(), + streaming=True, + submission_granularity="B", + consumption_granularity="B", + ) + it = gen.get_grouped_rollouts(request) + try: + await asyncio.wait_for(anext(it), timeout=10) + await asyncio.wait_for(anext(it), timeout=10) + await _flush() + gate = gen._active_pipeline.gate + # Batch 0 fully yielded but the consumer hasn't come back yet: its + # single batch slot is still held (a per-group release here would + # show release_calls == 2 and prepared == 4). + assert gate.release_calls == 0 + assert gen.prepare_group_rollout_calls == 2 + await asyncio.wait_for(anext(it), timeout=10) + await _flush() + assert gate.release_calls == 1 + assert gen.prepare_group_rollout_calls == 4 + finally: + await it.aclose() + + class TestRewardRollouts: @pytest.mark.asyncio async def test_get_reward_rollouts_matches_per_rollout_composition(self): @@ -348,3 +454,172 @@ def test_multi_env_distribution_requires_num_groups_above_one( assert min(agent_groups) == 0 assert all(slots == 0 for slots in agent_slots) assert np.gcd.reduce(agent_slots) == 0 + + +def make_response(epochs, prompt_length, total_len, content="resp", finish_reason="stop"): + return InferenceResponse( + response=LLMChatMessage(role="assistant", content=content), + raw_text=content, + token_ids=list(range(total_len)), + prompt_length=prompt_length, + logprobs=[0.0] * (total_len - prompt_length), + finish_reason=finish_reason, + policy_epoch=epochs, + kv_cache_epoch=epochs, + num_evictions=0, + ) + + +# Conversation length -> response spec: length 1 is the first turn (the bare prompt), length 3 +# the second (assistant reply + observation appended). +TWO_TURN_SCRIPT = { + 1: dict(epochs=[(0, 5)], prompt_length=3, total_len=7, content="a0"), + 3: dict(epochs=[(0, 5)], prompt_length=6, total_len=11, content="a1"), +} + +# Both two-turn termination modes (env-signaled done, max_turns exhausted) must produce this +# identical episode; only the env-consultation trace (observation_turns) differs per case. +TWO_TURN_EXPECTED = dict( + seen_roles=[["user"], ["user", "assistant", "user"]], + reward_conv=[("user", "hello"), ("assistant", "a0"), ("user", "obs0"), ("assistant", "a1")], + rewarded=[("a1", "stop")], + genmask_sums=[4, 5], + policy_epoch=[[(0, 5)], [(0, 5)]], +) + + +class ScriptedInterface(ReturnsTokens, ReturnsRaw): + """Inference stub whose reply is a pure function of the request: the conversation length + maps to a response spec, so it stays deterministic under pipeline concurrency.""" + + by_prompt_length: dict = Field(default_factory=dict) + seen_conversations: list = Field(default_factory=list) + + async def agenerate(self, request): + self.seen_conversations.append(list(request.prompt)) + return make_response(**self.by_prompt_length[len(request.prompt)]) + + +class EpisodeAgent(RewardOnlyAgent): + """Configurable multi-turn agent. + + `done_at_turn` controls when get_observation signals done: at every turn >= done_at_turn + it returns (None, True); None means it never signals done, so the episode ends only by + exhausting max_turns. Records get_reward calls and the conversation get_trajectory_reward saw. + """ + + env_id: str = "test" + max_turns: int = 1 + done_at_turn: int | None = None + rewarded: list = Field(default_factory=list) + reward_conversation: list = Field(default_factory=list) + observation_turns: list = Field(default_factory=list) + + async def get_prompt(self, validation): + return "hello", {"problem_id": "p0"} + + async def get_observation(self, turn_idx, response, conversation, golden): + self.observation_turns.append(turn_idx) + if self.done_at_turn is not None and turn_idx >= self.done_at_turn: + return None, True + return f"obs{turn_idx}", False + + async def get_reward(self, response, golden, finish_reason): + self.rewarded.append((response, finish_reason)) + return 1.5 + + async def get_trajectory_reward(self, responses, conversation, golden): + self.reward_conversation.extend(conversation) + return await super().get_trajectory_reward(responses, conversation, golden) + + +class TestMultiTurnEpisode: + + @pytest.mark.parametrize("driver", ["reward_rollouts", "pipeline"]) + @pytest.mark.parametrize( + "max_turns, done_at_turn, scripted, expected", + [ + # Single turn: get_observation is never consulted (no continuation is possible). + pytest.param( + 1, + None, + {1: dict(epochs=[(0, 7)], prompt_length=2, total_len=6, content="only")}, + dict( + seen_roles=[["user"]], + reward_conv=[("user", "hello"), ("assistant", "only")], + rewarded=[("only", "stop")], + genmask_sums=[4], + policy_epoch=[[(0, 7)]], + observation_turns=[], + ), + id="single_turn", + ), + # Multi-turn ended by the environment: turn 0 yields an observation, turn 1 is done. + pytest.param( + 3, + 1, + TWO_TURN_SCRIPT, + dict(TWO_TURN_EXPECTED, observation_turns=[0, 1]), + id="multi_turn_env_done", + ), + # Ended by exhausting max_turns instead (env never signals done): the same episode, + # except get_observation must not run for the final allowed turn. + pytest.param( + 2, + None, + TWO_TURN_SCRIPT, + dict(TWO_TURN_EXPECTED, observation_turns=[0]), + id="multi_turn_max_turns_exhausted", + ), + ], + ) + @pytest.mark.asyncio + async def test_run_episode(self, driver, max_turns, done_at_turn, scripted, expected): + """Episodes grow the conversation each turn and collapse into one per-turn rollout, + identically through get_reward_rollouts and through the real _RolloutPipeline + (get_grouped_rollouts) -- the latter proving run_episode runs in the infer stage.""" + iface = ScriptedInterface(by_prompt_length=scripted) + agent = EpisodeAgent(max_turns=max_turns, done_at_turn=done_at_turn) + + if driver == "reward_rollouts": + rollouts = await agent.get_reward_rollouts( + RolloutRequest(num_rollouts=1, inference_interface=iface) + ) + else: + groups = [] + + async def _drain(): + async for group in agent.get_grouped_rollouts( + GroupedRolloutRequest( + num_groups=1, rollouts_per_group=1, inference_interface=iface + ) + ): + groups.append(group) + + # Bounded so a wedged pipeline fails fast instead of hanging. + await asyncio.wait_for(_drain(), timeout=5.0) + (group,) = groups + rollouts = group.rollouts + (rollout,) = rollouts + + assert isinstance(rollout, TokenRollout) + assert rollout.reward == 1.5 + assert rollout.problem_id == "p0" + # One trajectory entry per generated turn. + assert len(rollout.trajectory) == len(expected["genmask_sums"]) + # Each turn's inference request = prior conversation (reply + observation appended). + assert [[m.role for m in conv] for conv in iface.seen_conversations] == expected[ + "seen_roles" + ] + # Default trajectory reward scores only the final response. + assert agent.rewarded == expected["rewarded"] + # Per-turn generation masks cover exactly each turn's generated tokens. + assert [sum(mask) for mask in rollout.generation_mask] == expected["genmask_sums"] + # Per-turn (engine-frame) staleness nesting is preserved. + assert rollout.policy_epoch == expected["policy_epoch"] + assert rollout.kv_cache_epoch == expected["policy_epoch"] + # get_observation is consulted only when another generation is still possible -- never on + # the final allowed turn. + assert agent.observation_turns == expected["observation_turns"] + # get_trajectory_reward sees the full dialogue, ending on the final reply exactly once. + assert [(m.role, m.content) for m in agent.reward_conversation] == expected["reward_conv"] diff --git a/tests/unit_tests/run_ci_test.sh b/tests/unit_tests/run_ci_test.sh index a51f7a21449..817dde579ed 100755 --- a/tests/unit_tests/run_ci_test.sh +++ b/tests/unit_tests/run_ci_test.sh @@ -145,9 +145,6 @@ DISTRIBUTED_ARGS=( --redirects "3" ) -# Reduce memory usage by NCCL -export NCCL_MAX_NCHANNELS=1 -export NCCL_NVLS_ENABLE=0 export ONE_LOGGER_JOB_CATEGORY=test # Run a pytest command. On marker-driven platforms a bucket can legitimately diff --git a/tests/unit_tests/ssm/ops/test_ssd_combined.py b/tests/unit_tests/ssm/ops/test_ssd_combined.py index b5ef14f7a79..6ca8f466023 100644 --- a/tests/unit_tests/ssm/ops/test_ssd_combined.py +++ b/tests/unit_tests/ssm/ops/test_ssd_combined.py @@ -5,6 +5,10 @@ import torch try: + from megatron.core.ssm.ops.intermediate_extraction import ( + scatter_intermediate_conv, + scatter_intermediate_ssm, + ) from megatron.core.ssm.ops.ssd_combined import is_int_pow_2, mamba_chunk_scan_combined_varlen HAVE_SSD_OPS = True @@ -158,7 +162,15 @@ def test_mamba_chunk_scan_combined_varlen_single_sequence(self): @unittest.skipIf(not HAVE_SSD_OPS, "SSD ops (Triton 3+) not available") @unittest.skipIf(not torch.cuda.is_available(), "CUDA required for SSD ops") class TestIntermediateStateExtraction(unittest.TestCase): - """Tests for intermediate_chunk_indices parameter.""" + """Tests for the ``return_raw_states=True`` contract of the chunk scan. + + Intermediate extraction was refactored: the kernel no longer gathers + requested chunks internally (old ``intermediate_chunk_indices`` / + ``return_intermediate_states`` kwargs). Instead the scan returns the full + ``(nchunks, ...)`` raw state tensor via ``return_raw_states=True``, and the + caller extracts from it with the fused ``scatter_intermediate_ssm`` kernel + (covered directly in TestScatterIntermediateKernels below). + """ def setUp(self): torch.manual_seed(42) @@ -181,8 +193,8 @@ def _make_inputs(self, seqlen): ) return x, dt, A, B, C, out - def test_intermediate_states_shape_and_no_nan(self): - """1 sequence, 4 chunks. Request intermediates at chunks [0, 1, 2].""" + def test_raw_states_shape_and_no_nan(self): + """1 sequence, 4 chunks. return_raw_states yields (final, all-chunk) states.""" seqlen = 64 # 4 chunks of 16 nchunks = seqlen // self.chunk_size x, dt, A, B, C, out = self._make_inputs(seqlen) @@ -191,7 +203,6 @@ def test_intermediate_states_shape_and_no_nan(self): ) last_chunk_indices = torch.tensor([nchunks - 1], dtype=torch.int64, device=self.device) seq_idx = torch.zeros(nchunks, dtype=torch.int32, device=self.device) - intermediate_chunk_indices = torch.tensor([0, 1, 2], dtype=torch.int64, device=self.device) result = mamba_chunk_scan_combined_varlen( x=x, @@ -204,18 +215,22 @@ def test_intermediate_states_shape_and_no_nan(self): last_chunk_indices=last_chunk_indices, seq_idx=seq_idx, out=out, - intermediate_chunk_indices=intermediate_chunk_indices, + return_raw_states=True, ) self.assertIsInstance(result, tuple) - final_states, intermediate_states = result + final_states, raw_states = result self.assertEqual(final_states.shape, (1, self.nheads, self.headdim, self.dstate)) - self.assertEqual(intermediate_states.shape, (3, self.nheads, self.headdim, self.dstate)) + # raw_states is the full per-chunk boundary state tensor. + self.assertEqual(raw_states.shape, (nchunks, self.nheads, self.headdim, self.dstate)) self.assertFalse(torch.isnan(final_states).any()) - self.assertFalse(torch.isnan(intermediate_states).any()) + self.assertFalse(torch.isnan(raw_states).any()) + # The final state is exactly the last chunk's raw state. + torch.testing.assert_close(final_states[0], raw_states[nchunks - 1]) - def test_intermediate_states_match_full_states(self): - """Intermediate states should match corresponding entries from full states.""" + def test_scatter_from_scan_raw_states(self): + """End-to-end: extract requested chunks from scan raw_states via the fused + scatter kernel and compare against the reference gather.""" seqlen = 64 # 4 chunks nchunks = seqlen // self.chunk_size x, dt, A, B, C, out = self._make_inputs(seqlen) @@ -225,9 +240,7 @@ def test_intermediate_states_match_full_states(self): last_chunk_indices = torch.tensor([nchunks - 1], dtype=torch.int64, device=self.device) seq_idx = torch.zeros(nchunks, dtype=torch.int32, device=self.device) - # Run with return_intermediate_states=True to get all states - out1 = torch.empty_like(out) - all_states = mamba_chunk_scan_combined_varlen( + final_states, raw_states = mamba_chunk_scan_combined_varlen( x=x, dt=dt, A=A, @@ -237,63 +250,42 @@ def test_intermediate_states_match_full_states(self): cu_chunk_seqlens=cu_chunk_seqlens, last_chunk_indices=last_chunk_indices, seq_idx=seq_idx, - out=out1, - return_intermediate_states=True, + out=out, + return_raw_states=True, ) + raw_states = raw_states.contiguous() - # Run with intermediate_chunk_indices indices = [0, 1, 2] - intermediate_chunk_indices = torch.tensor(indices, dtype=torch.int64, device=self.device) - out2 = torch.empty_like(out) - final_states, intermediate_states = mamba_chunk_scan_combined_varlen( - x=x, - dt=dt, - A=A, - B=B, - C=C, - chunk_size=self.chunk_size, - cu_chunk_seqlens=cu_chunk_seqlens, - last_chunk_indices=last_chunk_indices, - seq_idx=seq_idx, - out=out2, - intermediate_chunk_indices=intermediate_chunk_indices, + chunk_indices = torch.tensor(indices, dtype=torch.int64, device=self.device) + real_count_gpu = torch.tensor([len(indices)], dtype=torch.int32, device=self.device) + scratch = torch.empty( + len(indices), + self.nheads, + self.headdim, + self.dstate, + device=self.device, + dtype=raw_states.dtype, ) + scatter_intermediate_ssm(raw_states, chunk_indices, real_count_gpu, scratch) - # Intermediate states should match the corresponding all_states entries - for i, chunk_idx in enumerate(indices): - torch.testing.assert_close( - intermediate_states[i], - all_states[chunk_idx], - msg=f"intermediate state at index {i} (chunk {chunk_idx}) does not match", - ) + torch.testing.assert_close(scratch, raw_states[chunk_indices]) - # Final state should match last chunk - torch.testing.assert_close(final_states[0], all_states[nchunks - 1]) - - def test_intermediate_states_multi_sequence(self): - """2 packed sequences, verify intermediate extraction across sequence boundaries.""" + def test_scatter_from_scan_raw_states_multi_sequence(self): + """2 packed sequences: extract chunks that straddle a sequence boundary.""" seq1_len = 32 # 2 chunks seq2_len = 48 # 3 chunks total_len = seq1_len + seq2_len x, dt, A, B, C, out = self._make_inputs(total_len) - # cu_chunk_seqlens: seq1 has chunks at [0, 16, 32], seq2 at [32, 48, 64, 80] boundaries = list(range(0, seq1_len + 1, self.chunk_size)) + list( range(seq1_len + self.chunk_size, total_len + 1, self.chunk_size) ) cu_chunk_seqlens = torch.tensor(boundaries, dtype=torch.int32, device=self.device) nchunks = len(boundaries) - 1 # 5 chunks total - # Last chunk for seq1 is chunk 1, for seq2 is chunk 4 last_chunk_indices = torch.tensor([1, 4], dtype=torch.int64, device=self.device) - # seq_idx: [0, 0, 1, 1, 1] seq_idx = torch.tensor([0, 0, 1, 1, 1], dtype=torch.int32, device=self.device) - # Request chunk 0 from seq1 and chunks 2, 3 from seq2 - intermediate_chunk_indices = torch.tensor([0, 2, 3], dtype=torch.int64, device=self.device) - - # Also get full states for comparison - out_full = torch.empty_like(out) - all_states = mamba_chunk_scan_combined_varlen( + final_states, raw_states = mamba_chunk_scan_combined_varlen( x=x, dt=dt, A=A, @@ -303,34 +295,31 @@ def test_intermediate_states_multi_sequence(self): cu_chunk_seqlens=cu_chunk_seqlens, last_chunk_indices=last_chunk_indices, seq_idx=seq_idx, - out=out_full, - return_intermediate_states=True, - ) - - out2 = torch.empty_like(out) - final_states, intermediate_states = mamba_chunk_scan_combined_varlen( - x=x, - dt=dt, - A=A, - B=B, - C=C, - chunk_size=self.chunk_size, - cu_chunk_seqlens=cu_chunk_seqlens, - last_chunk_indices=last_chunk_indices, - seq_idx=seq_idx, - out=out2, - intermediate_chunk_indices=intermediate_chunk_indices, + out=out, + return_raw_states=True, ) - + raw_states = raw_states.contiguous() self.assertEqual(final_states.shape, (2, self.nheads, self.headdim, self.dstate)) - self.assertEqual(intermediate_states.shape, (3, self.nheads, self.headdim, self.dstate)) + self.assertEqual(raw_states.shape, (nchunks, self.nheads, self.headdim, self.dstate)) + + # Request chunk 0 from seq1 and chunks 2, 3 from seq2. + indices = [0, 2, 3] + chunk_indices = torch.tensor(indices, dtype=torch.int64, device=self.device) + real_count_gpu = torch.tensor([len(indices)], dtype=torch.int32, device=self.device) + scratch = torch.empty( + len(indices), + self.nheads, + self.headdim, + self.dstate, + device=self.device, + dtype=raw_states.dtype, + ) + scatter_intermediate_ssm(raw_states, chunk_indices, real_count_gpu, scratch) - # Verify intermediate states match full states - for i, chunk_idx in enumerate([0, 2, 3]): - torch.testing.assert_close(intermediate_states[i], all_states[chunk_idx]) + torch.testing.assert_close(scratch, raw_states[chunk_indices]) - def test_no_intermediate_returns_tensor(self): - """Without intermediate_chunk_indices, result should be a plain tensor.""" + def test_no_raw_states_returns_tensor(self): + """Without return_raw_states, result should be a plain final-state tensor.""" seqlen = 32 nchunks = seqlen // self.chunk_size x, dt, A, B, C, out = self._make_inputs(seqlen) @@ -357,5 +346,123 @@ def test_no_intermediate_returns_tensor(self): self.assertEqual(result.shape, (1, self.nheads, self.headdim, self.dstate)) +@unittest.skipIf(not HAVE_SSD_OPS, "SSD ops (Triton 3+) not available") +@unittest.skipIf(not torch.cuda.is_available(), "CUDA required for SSD ops") +class TestScatterIntermediateKernels(unittest.TestCase): + """Direct tests for the fused gather+scatter extraction kernels. + + These are the sole correctness coverage for scatter_intermediate_ssm / + scatter_intermediate_conv: equivalence vs. a reference gather, real_count + gating (padded slots left untouched), and the sub-d_conv clamp. + """ + + def setUp(self): + torch.manual_seed(0) + self.device = torch.device("cuda") + self.nheads = 4 + self.headdim = 16 + self.dstate = 8 + + def _ssm_states(self, num_chunks): + return torch.randn(num_chunks, self.nheads, self.headdim, self.dstate, device=self.device) + + @staticmethod + def _ref_conv(src, abs_positions, real_count, d_conv): + """Reference for scatter_intermediate_conv: for each meaningful slot, gather + the window [pos - d_conv, pos) (clamped into [0, seq_len-1]) and store it + transposed as out[slot, c, j].""" + _, seq_len, conv_dim = src.shape + max_count = abs_positions.shape[0] + out = torch.zeros(max_count, conv_dim, d_conv, device=src.device, dtype=src.dtype) + for slot in range(real_count): + pos = int(abs_positions[slot].item()) + for j in range(d_conv): + p = max(0, min(pos - d_conv + j, seq_len - 1)) + out[slot, :, j] = src[0, p, :] + return out + + def test_scatter_ssm_matches_reference(self): + """Fused gather matches the dense states[chunk_indices] reference.""" + states = self._ssm_states(num_chunks=6) + indices = [4, 0, 2] + chunk_indices = torch.tensor(indices, dtype=torch.int64, device=self.device) + real_count_gpu = torch.tensor([len(indices)], dtype=torch.int32, device=self.device) + out = torch.empty(len(indices), self.nheads, self.headdim, self.dstate, device=self.device) + + scatter_intermediate_ssm(states, chunk_indices, real_count_gpu, out) + + torch.testing.assert_close(out, states[chunk_indices]) + + def test_scatter_ssm_real_count_gating(self): + """Slots >= real_count are never written (padded scratch left untouched).""" + states = self._ssm_states(num_chunks=6) + max_count, real_count = 5, 3 + # Trailing indices are valid but must NOT be gathered (gated out). + chunk_indices = torch.tensor([4, 0, 2, 1, 5], dtype=torch.int64, device=self.device) + real_count_gpu = torch.tensor([real_count], dtype=torch.int32, device=self.device) + sentinel = 12345.0 + out = torch.full( + (max_count, self.nheads, self.headdim, self.dstate), sentinel, device=self.device + ) + + scatter_intermediate_ssm(states, chunk_indices, real_count_gpu, out) + + # First real_count slots gathered... + torch.testing.assert_close(out[:real_count], states[chunk_indices[:real_count]]) + # ...trailing slots left at the sentinel (no HBM write). + self.assertTrue(torch.all(out[real_count:] == sentinel)) + + def test_scatter_conv_matches_reference(self): + """Fused conv-window gather matches the transposed reference (positions in range).""" + d_conv = 4 + seq_len, conv_dim = 32, 12 + src = torch.randn(1, seq_len, conv_dim, device=self.device) + abs_positions = torch.tensor([10, 20, 5], dtype=torch.int32, device=self.device) + real_count_gpu = torch.tensor([3], dtype=torch.int32, device=self.device) + out = torch.empty(3, conv_dim, d_conv, device=self.device) + + scatter_intermediate_conv(src, abs_positions, real_count_gpu, out, d_conv) + + ref = self._ref_conv(src, abs_positions, real_count=3, d_conv=d_conv) + torch.testing.assert_close(out, ref) + + def test_scatter_conv_sub_dconv_clamp(self): + """A window whose start falls below token 0 clamps into range (reads token 0).""" + d_conv = 4 + seq_len, conv_dim = 32, 12 + src = torch.randn(1, seq_len, conv_dim, device=self.device) + # slot 0: pos=2 < d_conv -> window [-2, -1, 0, 1] clamps to [0, 0, 0, 1]. + abs_positions = torch.tensor([2, 20], dtype=torch.int32, device=self.device) + real_count_gpu = torch.tensor([2], dtype=torch.int32, device=self.device) + out = torch.empty(2, conv_dim, d_conv, device=self.device) + + scatter_intermediate_conv(src, abs_positions, real_count_gpu, out, d_conv) + + ref = self._ref_conv(src, abs_positions, real_count=2, d_conv=d_conv) + torch.testing.assert_close(out, ref) + # Explicitly: the three out-of-range positions all clamp to token 0. + torch.testing.assert_close(out[0, :, 0], src[0, 0, :]) + torch.testing.assert_close(out[0, :, 1], src[0, 0, :]) + torch.testing.assert_close(out[0, :, 2], src[0, 0, :]) + torch.testing.assert_close(out[0, :, 3], src[0, 1, :]) + + def test_scatter_conv_real_count_gating(self): + """Slots >= real_count are never written by the conv kernel either.""" + d_conv = 4 + seq_len, conv_dim = 32, 12 + src = torch.randn(1, seq_len, conv_dim, device=self.device) + max_count, real_count = 4, 2 + abs_positions = torch.tensor([10, 20, 15, 25], dtype=torch.int32, device=self.device) + real_count_gpu = torch.tensor([real_count], dtype=torch.int32, device=self.device) + sentinel = -999.0 + out = torch.full((max_count, conv_dim, d_conv), sentinel, device=self.device) + + scatter_intermediate_conv(src, abs_positions, real_count_gpu, out, d_conv) + + ref = self._ref_conv(src, abs_positions, real_count=real_count, d_conv=d_conv) + torch.testing.assert_close(out[:real_count], ref[:real_count]) + self.assertTrue(torch.all(out[real_count:] == sentinel)) + + if __name__ == "__main__": unittest.main() diff --git a/tests/unit_tests/ssm/test_gated_delta_net.py b/tests/unit_tests/ssm/test_gated_delta_net.py index 06e4c136d66..940196ccb55 100644 --- a/tests/unit_tests/ssm/test_gated_delta_net.py +++ b/tests/unit_tests/ssm/test_gated_delta_net.py @@ -1,5 +1,6 @@ # Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +import copy import os from functools import partial from unittest import mock @@ -9,9 +10,6 @@ import torch.nn.functional as F from megatron.core import parallel_state -from megatron.core.models.common.embeddings.rope_utils import ( - get_pos_emb_on_this_cp_rank as get_tensor_on_this_cp_rank, -) from megatron.core.models.gpt.experimental_attention_variant_module_specs import ( get_experimental_attention_variant_module_spec, get_transformer_block_with_experimental_attention_variant_spec, @@ -19,25 +17,16 @@ from megatron.core.models.gpt.gpt_model import GPTModel from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.process_groups_config import ProcessGroupCollection -from megatron.core.ssm.gated_delta_net import ( - GatedDeltaNet, +from megatron.core.ssm.gated_delta_net import GatedDeltaNet +from megatron.core.ssm.gated_delta_net.common import ( _build_head_perm_for_split_sections, _build_thd_cp_a2a_perm, tensor_a2a_cp2hp, tensor_a2a_hp2cp, + torch_chunk_gated_delta_rule, ) from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import TransformerConfig -from megatron.core.utils import unwrap_model -from megatron.training.arguments import parse_args -from megatron.training.checkpointing import load_checkpoint, save_checkpoint -from megatron.training.global_vars import set_args -from megatron.training.training import get_model -from tests.unit_tests.dist_checkpointing import ( - TempNamedDir, - init_basic_mock_args, - init_checkpointing_mock_args, -) from tests.unit_tests.test_utilities import Utils from tests.unit_tests.transformer.test_attention import _test_parallel_attention_correctness from tests.unit_tests.transformer.test_multi_latent_attention import ( @@ -385,6 +374,78 @@ def test_gpu_forward_rejects_sbhd_conv_padding(self): with pytest.raises(ValueError, match=expected_error): gdn(hidden_states, None) + def test_deterministic_mode(self): + tp_group = parallel_state.get_tensor_model_parallel_group() + cp_group = parallel_state.get_context_parallel_group() + pg_collection = ProcessGroupCollection(tp=tp_group, cp=cp_group) + + det_config = copy.deepcopy(self.transformer_config) + det_config.deterministic_mode = True + + gdn_submodules = get_experimental_attention_variant_module_spec( + config=det_config + ).submodules + + model_parallel_cuda_manual_seed(42) + torch.manual_seed(42) + gdn = ( + GatedDeltaNet( + det_config, + submodules=gdn_submodules, + layer_number=1, + bias=False, + conv_bias=False, + conv_init=1.0, + use_qk_l2norm=True, + A_init_range=(1, 16), + pg_collection=pg_collection, + ) + .cuda() + .bfloat16() + ) + + # deterministic_mode must select the torch-native kernel, not FLA. + assert gdn.gated_delta_rule is torch_chunk_gated_delta_rule + + micro_batch_size = 2 + seq_length = 64 + torch.manual_seed(0) + base_input = torch.randn( + (seq_length // self.sp_size // self.cp_size, micro_batch_size, gdn.config.hidden_size), + device=torch.cuda.current_device(), + dtype=torch.bfloat16, + ) + + def run(): + hidden_states = base_input.clone().requires_grad_(True) + output, _ = gdn(hidden_states, None) + output.float().sum().backward() + grads = { + name: param.grad.detach().clone() + for name, param in gdn.named_parameters() + if param.grad is not None + } + gdn.zero_grad(set_to_none=True) + return output.detach().clone(), grads, hidden_states.grad.detach().clone() + + out1, grads1, input_grad1 = run() + out2, grads2, input_grad2 = run() + + rank = torch.distributed.get_rank() + assert torch.equal(out1, out2), f"Output not reproducible ({rank=})" + assert torch.equal(input_grad1, input_grad2), f"Input grad not reproducible ({rank=})" + assert set(grads1.keys()) == set(grads2.keys()) + for name in grads1: + assert torch.equal( + grads1[name], grads2[name] + ), f"Grad not reproducible for {name} ({rank=})" + + def test_module_construction(self): + gdn = self.gdn + assert gdn.in_proj_dim == 2 * gdn.qk_dim + 2 * gdn.v_dim + 2 * gdn.num_value_heads + assert gdn.A_log.shape == (gdn.num_value_heads // self.tp_size,) + assert gdn.dt_bias.shape == (gdn.num_value_heads // self.tp_size,) + def test_jit_compiled_helpers(self): import torch._dynamo diff --git a/tests/unit_tests/ssm/test_hybrid_block.py b/tests/unit_tests/ssm/test_hybrid_block.py index 48946d0b42f..954dd486876 100644 --- a/tests/unit_tests/ssm/test_hybrid_block.py +++ b/tests/unit_tests/ssm/test_hybrid_block.py @@ -3,6 +3,7 @@ import pytest import torch +from megatron.core.extensions.transformer_engine import TEDotProductAttention from megatron.core.models.hybrid.hybrid_block import HybridStack, HyperConnectionHybridLayer from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols, validate_segment_layers from megatron.core.models.hybrid.hybrid_layer_specs import hybrid_stack_spec @@ -17,6 +18,7 @@ ) from megatron.core.transformer.experimental_attention_variant.dsa import DSAttention from megatron.core.transformer.mlp import MLP +from megatron.core.transformer.multi_latent_attention import MLASelfAttention from megatron.core.transformer.transformer_config import MLATransformerConfig from megatron.core.transformer.transformer_layer import TransformerLayer from tests.unit_tests.test_utilities import Utils @@ -116,6 +118,35 @@ def get_dsa_mamba_block(self, layer_pattern, enable_hyper_connections=False): pg_collection=self.get_pg_collection(), ) + def get_mla_hybrid_block(self, layer_pattern): + layer_type_list = validate_segment_layers(layer_pattern) + transformer_config = MLATransformerConfig( + hidden_size=256, # The Mamba layer places several constraints on this + # Need to specify num_attention_heads and num_layers or TransformerConfig + # will generate errors. + num_layers=len(layer_type_list), + num_attention_heads=16, + use_cpu_initialization=True, + bf16=True, + params_dtype=torch.bfloat16, + q_lora_rank=64, + kv_lora_rank=64, + qk_head_dim=64, + qk_pos_emb_head_dim=32, + v_head_dim=64, + rope_type='rope', + rotary_base=10000, + rotary_percent=1.0, + ) + modules = hybrid_stack_spec.submodules + return HybridStack( + transformer_config, + modules, + layer_type_list=layer_type_list, + pp_layer_offset=0, + pg_collection=self.get_pg_collection(), + ) + def teardown_method(self, method): Utils.destroy_model_parallel() @@ -399,7 +430,7 @@ def test_hyper_connection_pipeline_boundary_shapes(self): assert output.shape == (sequence_length, micro_batch_size, transformer_config.hidden_size) def test_invalid_layer_types_cause_failure(self): - invalid_symbol = '+' + invalid_symbol = 'X' assert invalid_symbol not in Symbols.VALID_LAYERS # sanity check. layer_pattern = Symbols.MAMBA + Symbols.ATTENTION + Symbols.MLP + invalid_symbol # validate_segment_layers() in hybrid_layer_allocation.py throws a ValueError. @@ -470,3 +501,21 @@ def test_mixed_attention_and_dsa_layer_types(self): layer_pattern = Symbols.MAMBA + Symbols.ATTENTION + Symbols.DS_ATTENTION + Symbols.MAMBA with pytest.raises(ValueError): block = self.get_dsa_mamba_block(layer_pattern) + + def test_mla_layer_types(self): + """+ symbol creates a TransformerLayer with MLASelfAttention but + standard (non-DSA) core attention.""" + layer_pattern = Symbols.MAMBA + Symbols.MLA + Symbols.MAMBA + block = self.get_mla_hybrid_block(layer_pattern) + layers = block.layers + assert isinstance(layers[0], MambaLayer) + assert isinstance(layers[1], TransformerLayer) + assert isinstance(layers[1].self_attention, MLASelfAttention) + assert isinstance(layers[1].self_attention.core_attention, TEDotProductAttention) + assert isinstance(layers[2], MambaLayer) + + def test_mixed_attention_and_mla_layer_types(self): + """* and + in the same block fail (same reason as * and D).""" + layer_pattern = Symbols.MAMBA + Symbols.ATTENTION + Symbols.MLA + Symbols.MAMBA + with pytest.raises(ValueError): + block = self.get_mla_hybrid_block(layer_pattern) diff --git a/tests/unit_tests/ssm/test_hybrid_layer_allocation.py b/tests/unit_tests/ssm/test_hybrid_layer_allocation.py index 2618d9cde50..de49033cdff 100644 --- a/tests/unit_tests/ssm/test_hybrid_layer_allocation.py +++ b/tests/unit_tests/ssm/test_hybrid_layer_allocation.py @@ -78,6 +78,7 @@ def test_valid_patterns(self): ("GGG*GGG*", ['G', 'G', 'G', '*', 'G', 'G', 'G', '*']), ("GEGEGE*E", ['G', 'E', 'G', 'E', 'G', 'E', '*', 'E']), ("MDMD", ['M', 'D', 'M', 'D']), + ("M+M+", ['M', '+', 'M', '+']), ] for pattern, expected in test_cases: result = validate_segment_layers(pattern) @@ -101,6 +102,11 @@ def test_invalid_symbols_cause_failure(self): with pytest.raises(ValueError): # Not allowed to have both standard Attention and MLA/DSA validate_segment_layers("MDM*-") + with pytest.raises(ValueError): + # Not allowed to have both standard Attention and MLA (same reason + # as DSA: * uses the model-level rotary_pos_emb while + uses MLA's + # own decoupled RoPE). + validate_segment_layers("M+M*-") def test_window_symbol(self): """'W' (sliding-window-only DSv4 attention) is a first-class MLA layer symbol.""" @@ -177,6 +183,8 @@ def test_main_pattern_only(self): ("GEGEGE*E", "GEGEGE*E"), ("MDMD", "MDMD"), ("DM", "DM"), + ("M+M+", "M+M+"), + ("+M", "+M"), ] for pattern, expected_main in test_cases: result = parse_hybrid_pattern(pattern) @@ -301,6 +309,8 @@ def test_complex_patterns(self): ("GEGEGE*E/GG/GG", "GEGEGE*E", "GG", 2), # DSA in main pattern with MTP ("MDMD/MD/MD", "MDMD", "MD", 2), + # MLA in main pattern with MTP + ("M+M+/M+/M+", "M+M+", "M+", 2), ] for pattern, expected_main, expected_mtp, expected_depths in test_cases: result = parse_hybrid_pattern(pattern) @@ -327,6 +337,7 @@ def test_simple_pattern(self): 'D': 0, 'G': 0, 'M': 2, + '+': 0, '-': 0, 'E': 0, } @@ -341,6 +352,7 @@ def test_all_layer_types(self): 'D': 0, 'G': 1, 'M': 1, + '+': 0, '-': 1, 'E': 1, } @@ -352,6 +364,19 @@ def test_all_layer_types(self): 'D': 1, 'G': 1, 'M': 1, + '+': 0, + '-': 1, + 'E': 1, + } + assert get_hybrid_layer_counts("MG+-E") == { + 'C': 0, + 'H': 0, + 'W': 0, + '*': 0, + 'D': 0, + 'G': 1, + 'M': 1, + '+': 1, '-': 1, 'E': 1, } @@ -366,6 +391,7 @@ def test_with_pipes(self): 'D': 0, 'G': 0, 'M': 2, + '+': 0, '-': 0, 'E': 0, } @@ -377,6 +403,7 @@ def test_with_pipes(self): 'D': 0, 'G': 0, 'M': 4, + '+': 0, '-': 4, 'E': 0, } @@ -391,6 +418,7 @@ def test_with_mtp(self): 'D': 0, 'G': 0, 'M': 6, + '+': 0, '-': 0, 'E': 0, } @@ -406,6 +434,7 @@ def test_with_pipes_and_mtp(self): 'D': 0, 'G': 0, 'M': 8, + '+': 0, '-': 4, 'E': 0, } @@ -419,6 +448,7 @@ def test_moe_pattern(self): 'D': 0, 'G': 0, 'M': 2, + '+': 0, '-': 0, 'E': 2, } @@ -433,6 +463,7 @@ def test_mtp_with_attention(self): 'D': 0, 'G': 0, 'M': 7, + '+': 0, '-': 0, 'E': 0, } @@ -446,6 +477,7 @@ def test_gdn_pattern(self): 'D': 0, 'G': 2, 'M': 2, + '+': 0, '-': 0, 'E': 0, } @@ -460,6 +492,7 @@ def test_gdn_hybrid_pattern(self): 'D': 0, 'G': 2, 'M': 1, + '+': 0, '-': 0, 'E': 0, } @@ -473,6 +506,21 @@ def test_dsa_pattern(self): 'D': 2, 'G': 0, 'M': 2, + '+': 0, + '-': 0, + 'E': 0, + } + + def test_mla_pattern(self): + assert get_hybrid_layer_counts("+M+M") == { + 'C': 0, + 'H': 0, + 'W': 0, + '*': 0, + 'D': 0, + 'G': 0, + 'M': 2, + '+': 2, '-': 0, 'E': 0, } @@ -486,6 +534,7 @@ def test_empty_pattern(self): 'D': 0, 'G': 0, 'M': 0, + '+': 0, '-': 0, 'E': 0, } @@ -814,3 +863,39 @@ def test_all_mamba(self): assert mamba_map == {0: 0, 1: 1, 2: 2} assert mlp_map == {} assert moe_map == {} + + def test_mla(self): + """+ (MLA) layers are mapped independently of other attention types.""" + maps = get_layer_maps_from_layer_type_list(["+", "M", "+", "M"]) + attention_map, dsa_map, mamba_map, mla_map, mlp_map, moe_map = operator.itemgetter( + Symbols.ATTENTION, + Symbols.DS_ATTENTION, + Symbols.MAMBA, + Symbols.MLA, + Symbols.MLP, + Symbols.MOE, + )(maps) + assert attention_map == {} + assert dsa_map == {} + assert mla_map == {0: 0, 2: 1} + assert mamba_map == {1: 0, 3: 1} + assert mlp_map == {} + assert moe_map == {} + + def test_mixed_dsa_and_mla(self): + """D and + can coexist (both are MLA-based and use decoupled RoPE).""" + maps = get_layer_maps_from_layer_type_list(["D", "+", "M", "-"]) + attention_map, dsa_map, mamba_map, mla_map, mlp_map, moe_map = operator.itemgetter( + Symbols.ATTENTION, + Symbols.DS_ATTENTION, + Symbols.MAMBA, + Symbols.MLA, + Symbols.MLP, + Symbols.MOE, + )(maps) + assert attention_map == {} + assert dsa_map == {0: 0} + assert mla_map == {1: 0} + assert mamba_map == {2: 0} + assert mlp_map == {3: 0} + assert moe_map == {} diff --git a/tests/unit_tests/ssm/test_mamba_layer.py b/tests/unit_tests/ssm/test_mamba_layer.py index 8d6e0ab8c91..aad23ce02d6 100644 --- a/tests/unit_tests/ssm/test_mamba_layer.py +++ b/tests/unit_tests/ssm/test_mamba_layer.py @@ -1,5 +1,7 @@ # Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. +from dataclasses import replace + import pytest import torch @@ -9,6 +11,7 @@ from megatron.core.ssm.mamba_layer import MambaLayer, MambaLayerSubmodules from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer import TransformerConfig +from megatron.core.transformer.torch_norm import WrappedTorchNorm from tests.unit_tests.test_utilities import Utils @@ -24,20 +27,24 @@ def setup_method(self, method): # will generate errors. num_layers=1, num_attention_heads=1, + layernorm_epsilon=1e-6, use_cpu_initialization=True, ) assert isinstance(hybrid_stack_spec.submodules, HybridStackSubmodules) assert isinstance(hybrid_stack_spec.submodules.mamba_layer.submodules, MambaLayerSubmodules) - pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp']) - self.layer = MambaLayer( - transformer_config, - hybrid_stack_spec.submodules.mamba_layer.submodules, - pg_collection=pg_collection, + # Use an explicit norm so the test can verify the configured epsilon. + mamba_submodules = replace( + hybrid_stack_spec.submodules.mamba_layer.submodules, norm=WrappedTorchNorm ) + pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['tp', 'cp']) + self.layer = MambaLayer(transformer_config, mamba_submodules, pg_collection=pg_collection) def teardown_method(self, method): Utils.destroy_model_parallel() + def test_configured_layernorm_epsilon(self): + assert self.layer.norm.eps == self.layer.config.layernorm_epsilon + def test_gpu_forward(self): layer = self.layer layer.cuda() diff --git a/tests/unit_tests/ssm/test_split_tensor_factory.py b/tests/unit_tests/ssm/test_split_tensor_factory.py index abb668e16a8..ab9fd434e08 100644 --- a/tests/unit_tests/ssm/test_split_tensor_factory.py +++ b/tests/unit_tests/ssm/test_split_tensor_factory.py @@ -7,7 +7,7 @@ import torch from megatron.core.dist_checkpointing import ShardedTensor -from megatron.core.ssm.gated_delta_net import ( +from megatron.core.ssm.gated_delta_net.common import ( _split_tensor_factory as gated_delta_split_tensor_factory, ) from megatron.core.ssm.mamba_mixer import _split_tensor_factory as mamba_split_tensor_factory diff --git a/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py b/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py index 44d6fa21178..a76746d7674 100644 --- a/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py +++ b/tests/unit_tests/tensor_parallel/test_tp_attrs_without_init.py @@ -8,6 +8,7 @@ RowParallelLinear, VocabParallelEmbedding, copy_tensor_model_parallel_attributes, + param_is_not_tensor_parallel_duplicate, ) from megatron.core.transformer.transformer_config import TransformerConfig from tests.unit_tests.test_utilities import Utils @@ -100,3 +101,60 @@ def test_copy_tensor_model_parallel_attributes_preserves_qkv_split_shapes(): assert destination.is_qkv is True assert destination.qkv_split_shapes == source.qkv_split_shapes + + +def test_non_allreduce_param_uses_expert_tp_group_for_duplicate_filter(): + class RankGroup: + def __init__(self, rank): + self._rank = rank + + def rank(self): + return self._rank + + param = torch.empty(1) + regular_tp_group = RankGroup(rank=1) + expert_tp_group = RankGroup(rank=0) + assert not param_is_not_tensor_parallel_duplicate(param, regular_tp_group) + + param.allreduce = False + assert param_is_not_tensor_parallel_duplicate( + param, tp_group=regular_tp_group, expert_tp_group=expert_tp_group + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") +def test_expert_linear_parameters_use_expert_topology_metadata(): + Utils.initialize_model_parallel(tensor_model_parallel_size=2, expert_tensor_parallel_size=1) + cfg = TransformerConfig( + num_layers=1, + hidden_size=8, + num_attention_heads=4, + tensor_model_parallel_size=2, + expert_tensor_parallel_size=1, + use_cpu_initialization=True, + perform_initialization=False, + ) + + column = ColumnParallelLinear( + input_size=8, + output_size=8, + init_method=cfg.init_method, + bias=True, + config=cfg, + gather_output=False, + skip_bias_add=False, + is_expert=True, + ) + row = RowParallelLinear( + input_size=8, + output_size=8, + init_method=cfg.init_method, + bias=True, + input_is_parallel=True, + config=cfg, + skip_bias_add=False, + is_expert=True, + ) + + for param in (*column.parameters(), *row.parameters()): + assert param.allreduce is False diff --git a/tests/unit_tests/test_checkpointing.py b/tests/unit_tests/test_checkpointing.py index 9b62c26d674..cf1b6a76539 100644 --- a/tests/unit_tests/test_checkpointing.py +++ b/tests/unit_tests/test_checkpointing.py @@ -25,6 +25,7 @@ _load_base_checkpoint, get_checkpoint_tracker_filename, load_checkpoint, + maybe_save_dataloader_state, read_metadata, save_checkpoint, ) @@ -74,6 +75,80 @@ def sharded_state_dict(self, *args, metadata: Optional[dict] = None, **kwargs): return self.state_dict() +def test_maybe_save_dataloader_state_uses_explicit_process_groups(tmp_path): + """Dataloader checkpoints use the supplied module groups and canonical model-parallel path.""" + groups = { + "tp": SimpleNamespace(rank=0, size=2), + "pp": SimpleNamespace(rank=0, size=2), + "dp": SimpleNamespace(rank=3, size=4), + } + barriers = [] + saved = [] + iterator = SimpleNamespace( + iterable=SimpleNamespace(save_state=lambda: {"global_sequence_id": 16}) + ) + + with ( + mock.patch( + "megatron.training.checkpointing.get_pg_rank", side_effect=lambda group: group.rank + ), + mock.patch( + "megatron.training.checkpointing.get_pg_size", side_effect=lambda group: group.size + ), + mock.patch( + "megatron.training.checkpointing.torch.distributed.barrier", + side_effect=lambda group: barriers.append(group), + ), + mock.patch( + "megatron.training.checkpointing.torch.save", + side_effect=lambda state, path: saved.append((state, path)), + ), + ): + maybe_save_dataloader_state( + iterator, + 2, + tmp_path, + tp_group=groups["tp"], + pp_group=groups["pp"], + dp_group=groups["dp"], + ) + + assert barriers == [groups["dp"], groups["dp"]] + assert saved[0][0] == {"dataloader_state_dict": {"global_sequence_id": 16}} + assert saved[0][1] == str( + tmp_path / "iter_0000002" / "mp_rank_00_000" / "train_dataloader_dprank003.pt" + ) + + +def test_maybe_save_dataloader_state_skips_empty_state_after_barriers(tmp_path): + """Ranks without dataloader state participate in barriers but do not write a file.""" + group = SimpleNamespace(rank=0, size=1) + iterator = SimpleNamespace(iterable=SimpleNamespace(save_state=lambda: None)) + barriers = [] + + with ( + mock.patch( + "megatron.training.checkpointing.get_pg_rank", + side_effect=lambda process_group: process_group.rank, + ), + mock.patch( + "megatron.training.checkpointing.get_pg_size", + side_effect=lambda process_group: process_group.size, + ), + mock.patch( + "megatron.training.checkpointing.torch.distributed.barrier", + side_effect=lambda group: barriers.append(group), + ), + mock.patch("megatron.training.checkpointing.torch.save") as save, + ): + maybe_save_dataloader_state( + iterator, 2, tmp_path, tp_group=group, pp_group=group, dp_group=group + ) + + assert barriers == [group, group] + save.assert_not_called() + + def create_checkpoint(load_path, ckpt_format): """Setup a dummy checkpoint directory.""" iteration = 123 @@ -139,6 +214,7 @@ def create_ckpt_load_args(create_args): args.tensor_model_parallel_size = 1 args.pipeline_model_parallel_size = 1 args.ckpt_assume_constant_structure = False + args.stream_ckpt_dequant = True args.ckpt_fully_parallel_save = False args.ckpt_fully_parallel_load = False args.ckpt_load_validate_sharding_integrity = True diff --git a/tests/unit_tests/test_fp8_param.py b/tests/unit_tests/test_fp8_param.py index 69265906a4c..1179502eb62 100644 --- a/tests/unit_tests/test_fp8_param.py +++ b/tests/unit_tests/test_fp8_param.py @@ -26,7 +26,7 @@ set_args, set_global_variables, ) -from megatron.training.training import get_model, setup_model_and_optimizer +from megatron.training.training import force_param_sync, get_model, setup_model_and_optimizer from megatron.training.utils import get_device_arch_version from tests.unit_tests.test_utilities import Utils @@ -223,6 +223,9 @@ def _run_test_helper( **kwargs, ): """Test fp8_param with a small GPT model.""" + # Test-only knob: not a model arg, so pop before create_test_args (which asserts every + # kwarg is a real arg attribute). + save_at_steps_kw = kwargs.pop("save_at_steps", ()) args = self.create_test_args( tp_size, recipe, @@ -244,6 +247,9 @@ def _run_test_helper( Utils.initialize_model_parallel( tensor_model_parallel_size=tp_size, expert_model_parallel_size=args.expert_model_parallel_size, + # Enable GTP weight-remat when the test requested it (default 1 => no GTP, so + # non-GTP fp8 tests are unaffected). + gtp_remat_size=getattr(args, "gtp_weight_remat_size", 1), ) input_ids, labels, position_ids, attention_mask, loss_mask = self.get_batch( @@ -332,11 +338,27 @@ def _run_test_helper( loss_list = [] eval_loss_list = [] + # Optional: generate the sharded_state_dict (the checkpoint-save metadata path) at these + # steps to catch save side-effects on the live weights — a correct save must not perturb + # the subsequent training step (regression guard for GTP native-FP8 save corruption). + save_at_steps = set(save_at_steps_kw or ()) + for i in range(100): if not inference: gpt_model[0].zero_grad_buffer() optimizer.zero_grad() + if i in save_at_steps: + # Mirror production save_checkpoint_and_time: when the forward pre-hook is disabled + # for the save, a forced param-sync runs first. Passing the optimizer makes it copy + # the FP32 masters into the param buffer before the copy-back re-quantizes, so + # native-FP8 GTP shards are refreshed from masters (not stale grad scratch). + # Exercise it so the save-perturbation test is a real regression test for the + # post-save loss spike. + if should_disable_forward_pre_hook(args): + force_param_sync(gpt_model, optimizer=optimizer) + _ = gpt_model[0].sharded_state_dict() + # Capture CUDA graphs after warmup if helper is provided. # Hard coded cuda_graph_warmup_steps = 0. cuda_graph_warmup_steps = 0 diff --git a/tests/unit_tests/test_fp8_utils.py b/tests/unit_tests/test_fp8_utils.py index 5be17f03c9f..dc65d541455 100644 --- a/tests/unit_tests/test_fp8_utils.py +++ b/tests/unit_tests/test_fp8_utils.py @@ -7,8 +7,21 @@ import torch.nn as nn from megatron.core import fp8_utils +from megatron.training.utils import get_device_arch_version from tests.unit_tests.test_utilities import Utils +try: + import transformer_engine_torch as tex + from transformer_engine.pytorch.tensor.mxfp8_tensor import MXFP8Quantizer + + HAVE_MXFP8_TENSOR = True +except ImportError: + HAVE_MXFP8_TENSOR = False + +# MXFP8 needs Blackwell or newer. +mxfp8_available = HAVE_MXFP8_TENSOR and get_device_arch_version() >= 10 +reason_for_no_mxfp8 = "MXFP8 requires Transformer Engine and device arch >= 10" + class MockTELinear(nn.Module): """Mock TE Linear module for testing.""" @@ -130,3 +143,66 @@ def track_forward(x): # Verify output has original shape assert output.shape == (6, 2, 4096) # Back to original seq_len + + +@pytest.mark.skipif(not mxfp8_available, reason=reason_for_no_mxfp8) +class TestCopyTensorsToQuantizedParams: + """Cover the batched MXFP8 param copy-back used by _post_param_sync. + + ``copy_tensors_to_quantized_params`` bypasses ``copy_`` and calls the destination quantizer + directly, so the contract to protect is that it still writes exactly what the per-param + ``copy_tensor_to_quantized_param`` would have written. + """ + + SHAPES = [(1024, 512), (2048, 256), (512, 1024)] + + def _make_param(self, shape): + quantizer = MXFP8Quantizer(fp8_dtype=tex.DType.kFloat8E4M3, rowwise=True, columnwise=True) + tensor = quantizer.make_empty(shape, dtype=torch.bfloat16, device="cuda") + return torch.nn.Parameter(tensor, requires_grad=False) + + def _raw_buffers(self, param): + """The four buffers MXFP8 storage is made of, i.e. everything a cast writes.""" + data = param.data + return ( + data._rowwise_data, + data._rowwise_scale_inv, + data._columnwise_data, + data._columnwise_scale_inv, + ) + + def test_matches_per_param_copy(self): + """Batched copy-back is bitwise identical to copying one param at a time.""" + torch.manual_seed(0) + reference_params = [self._make_param(shape) for shape in self.SHAPES] + batched_params = [self._make_param(shape) for shape in self.SHAPES] + # Sources are flat slices, matching how _post_param_sync views the param buffer. + srcs = [ + torch.randn(shape, dtype=torch.bfloat16, device="cuda").view(-1) + for shape in self.SHAPES + ] + + for param, src in zip(reference_params, srcs): + fp8_utils.copy_tensor_to_quantized_param(param, src) + fp8_utils.copy_tensors_to_quantized_params(batched_params, srcs) + torch.cuda.synchronize() + + for reference, batched in zip(reference_params, batched_params): + for expected, actual in zip(self._raw_buffers(reference), self._raw_buffers(batched)): + assert torch.equal(expected, actual) + + def test_falls_back_without_quantizer(self): + """A destination with no quantizer of its own still gets written.""" + param = self._make_param(self.SHAPES[0]) + param.data._quantizer = None + src = torch.randn(self.SHAPES[0], dtype=torch.bfloat16, device="cuda").view(-1) + + fp8_utils.copy_tensors_to_quantized_params([param], [src]) + torch.cuda.synchronize() + + # A quantized copy of a non-zero source cannot be all zeros. + assert param.data._rowwise_data.any() + + def test_empty_input(self): + """No params is a no-op rather than an error.""" + fp8_utils.copy_tensors_to_quantized_params([], []) diff --git a/tests/unit_tests/test_frozen_ckpt_resume.py b/tests/unit_tests/test_frozen_ckpt_resume.py new file mode 100644 index 00000000000..3b9b5c44d63 --- /dev/null +++ b/tests/unit_tests/test_frozen_ckpt_resume.py @@ -0,0 +1,56 @@ +# Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. + +"""Unit tests for --freeze-all-layers auto-resume iteration reading. + +``read_frozen_resume_iteration`` reads the progress tracker +(``latest_checkpointed_iteration.txt``) that a frozen (--freeze-all-layers) run writes +to its --save dir, so an identical resubmitted job continues where it stopped instead +of restarting. Its main use today is offline-KD teacher-logit dumps. Pure file I/O +(no CUDA, no distributed init), so it runs on CPU. +""" + +from megatron.training.checkpointing import ( + get_checkpoint_tracker_filename, + read_frozen_resume_iteration, +) + + +def _write_tracker(dir_path, content): + with open(get_checkpoint_tracker_filename(str(dir_path)), "w") as f: + f.write(content) + + +def test_missing_tracker_is_fresh_dump(tmp_path): + """No tracker in the load dir -> start at iteration 0 (first run, --finetune-like).""" + assert read_frozen_resume_iteration(str(tmp_path)) == 0 + + +def test_none_load_dir_is_zero(): + """A None load dir (nothing to resume from) -> 0.""" + assert read_frozen_resume_iteration(None) == 0 + + +def test_reads_recorded_iteration(tmp_path): + """A tracker written by a prior dump is read back as the resume iteration.""" + _write_tracker(tmp_path, "1500") + assert read_frozen_resume_iteration(str(tmp_path)) == 1500 + + +def test_tolerates_trailing_whitespace(tmp_path): + """A trailing newline in the tracker is stripped (read_metadata semantics).""" + _write_tracker(tmp_path, "42\n") + assert read_frozen_resume_iteration(str(tmp_path)) == 42 + + +def test_release_tracker_is_zero(tmp_path): + """A 'release' tracker maps to iteration 0 (start from the beginning).""" + _write_tracker(tmp_path, "release") + assert read_frozen_resume_iteration(str(tmp_path)) == 0 + + +def test_advancing_progress_reads_latest(tmp_path): + """Overwriting the tracker (progress advancing across runs) reads the newest value.""" + _write_tracker(tmp_path, "100") + assert read_frozen_resume_iteration(str(tmp_path)) == 100 + _write_tracker(tmp_path, "250") + assert read_frozen_resume_iteration(str(tmp_path)) == 250 diff --git a/tests/unit_tests/test_lion_optimizer.py b/tests/unit_tests/test_lion_optimizer.py index b0df91073ed..be36f101bdd 100644 --- a/tests/unit_tests/test_lion_optimizer.py +++ b/tests/unit_tests/test_lion_optimizer.py @@ -19,6 +19,7 @@ _get_megatron_optimizer_based_on_param_groups, _get_param_groups, ) +from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer from megatron.core.optimizer.optimizer import FP32Optimizer requires_emerging_optimizers = pytest.mark.skipif( @@ -253,3 +254,53 @@ def test_megatron_lion_exact_match_with_standalone(self, lr, beta1, beta2, weigh rtol=0, msg=f"Step {step}, state '{key}': optimizer states differ", ) + + +class TestDistributedOptimizerStateKeys: + """Tests for DistributedOptimizer.optimizer_state_keys and _get_state_key_dtype. + + These tests use a mock to avoid needing a full distributed setup. + """ + + def _make_mock_distopt(self, optimizer_name): + """Create a minimal mock with just the config needed for optimizer_state_keys.""" + mock = object.__new__(DistributedOptimizer) + mock.config = OptimizerConfig(optimizer=optimizer_name, lr=1e-4) + return mock + + def test_adam_state_keys(self): + distopt = self._make_mock_distopt("adam") + assert distopt.optimizer_state_keys == ("exp_avg", "exp_avg_sq") + + def test_lion_state_keys(self): + distopt = self._make_mock_distopt("lion") + assert distopt.optimizer_state_keys == ("exp_avg",) + + def test_sgd_state_keys_defaults_to_adam(self): + distopt = self._make_mock_distopt("sgd") + assert distopt.optimizer_state_keys == ("exp_avg", "exp_avg_sq") + + def test_get_state_key_dtype_known_keys(self): + distopt = self._make_mock_distopt("adam") + assert distopt._get_state_key_dtype("exp_avg") == torch.float32 + assert distopt._get_state_key_dtype("exp_avg_sq") == torch.float32 + + def test_get_state_key_dtype_unknown_key(self): + distopt = self._make_mock_distopt("adam") + assert distopt._get_state_key_dtype("unknown_key") == torch.float32 + + def test_get_state_key_dtype_respects_config(self): + mock = object.__new__(DistributedOptimizer) + mock.config = SimpleNamespace(exp_avg_dtype=torch.bfloat16, exp_avg_sq_dtype=torch.float16) + assert mock._get_state_key_dtype("exp_avg") == torch.bfloat16 + assert mock._get_state_key_dtype("exp_avg_sq") == torch.float16 + + def test_muon_with_lion_scalar_optimizer(self): + mock = object.__new__(DistributedOptimizer) + mock.config = OptimizerConfig(optimizer="muon", lr=1e-4, muon_scalar_optimizer="lion") + assert mock.optimizer_state_keys == ("exp_avg",) + + def test_muon_with_adam_scalar_optimizer(self): + mock = object.__new__(DistributedOptimizer) + mock.config = OptimizerConfig(optimizer="muon", lr=1e-4, muon_scalar_optimizer="adam") + assert mock.optimizer_state_keys == ("exp_avg", "exp_avg_sq") diff --git a/tests/unit_tests/test_msc_utils.py b/tests/unit_tests/test_msc_utils.py new file mode 100644 index 00000000000..e30e74f928e --- /dev/null +++ b/tests/unit_tests/test_msc_utils.py @@ -0,0 +1,40 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. +import builtins +import os +from pathlib import Path + +import pytest + +from megatron.core.msc_utils import maybe_msc + + +def test_open_is_builtin_open(): + assert maybe_msc.open is builtins.open + + +def test_os_is_os_module(): + assert maybe_msc.os is os + + +def test_Path_is_pathlib_Path(): + assert maybe_msc.Path is Path + + +def test_unknown_attribute_raises(): + with pytest.raises(AttributeError): + getattr(maybe_msc, "this_attribute_does_not_exist_12345") + + +def test_path_isdir_delegates_to_os_path_isdir(monkeypatch): + called = {} + + def fake_isdir(p): + called['p'] = p + return True + + # monkeypatch the os.path.isdir used by the fallback path + monkeypatch.setattr(os.path, 'isdir', fake_isdir) + + result = maybe_msc.path_isdir('/tmp/some-path') + assert result is True + assert called['p'] == '/tmp/some-path' diff --git a/tests/unit_tests/test_optimizer.py b/tests/unit_tests/test_optimizer.py index 5b3e69c23b8..7d99f50c96a 100644 --- a/tests/unit_tests/test_optimizer.py +++ b/tests/unit_tests/test_optimizer.py @@ -23,6 +23,8 @@ get_megatron_optimizer, get_standard_config_overrides, ) +from megatron.core.optimizer.distrib_optimizer import DistributedOptimizer +from megatron.core.optimizer.optimizer import copy_optimizer_param_metadata from megatron.core.optimizer_param_scheduler import ParamGroupOverride from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.transformer import TransformerConfig @@ -69,6 +71,16 @@ def forward(self, x): return x +def test_copy_optimizer_param_metadata_preserves_allreduce(): + source = torch.empty(1) + destination = torch.empty_like(source) + source.allreduce = False + + copy_optimizer_param_metadata(destination, source) + + assert destination.allreduce is False + + @patch('torch.distributed.get_world_size', return_value=1) @patch( 'torch.distributed.all_gather_object', lambda output_list, obj: output_list.__setitem__(0, obj) diff --git a/tests/unit_tests/test_process_groups_config.py b/tests/unit_tests/test_process_groups_config.py index b49962b1a5a..a61936bd132 100644 --- a/tests/unit_tests/test_process_groups_config.py +++ b/tests/unit_tests/test_process_groups_config.py @@ -29,7 +29,7 @@ def test_transformer_process_groups(self, mocker): # Test attribute existence assert hasattr(model_pgs, 'tp') assert hasattr(model_pgs, 'pp') - assert not hasattr(model_pgs, 'cp') # Not set yet + assert model_pgs.cp is None # Not set yet def test_grad_comm_process_groups(self, mocker): """Test basic functionality of ProcessGroupCollection.""" @@ -47,7 +47,7 @@ def test_grad_comm_process_groups(self, mocker): # Test attribute existence assert hasattr(grad_pgs, 'dp') - assert not hasattr(grad_pgs, 'dp_cp') # Not set yet + assert grad_pgs.dp_cp is None # Not set yet def test_hierarchical_context_parallel_groups(self, mocker): """Test setting and accessing the hierarchical context parallel list.""" @@ -129,7 +129,7 @@ def test_default_initialization(self): assert hasattr(model_pgs, 'tp') assert hasattr(model_pgs, 'pp') assert hasattr(model_pgs, 'cp') - assert not hasattr(model_pgs, 'dp') + assert model_pgs.dp is None # Not requested, so not set # Test that an error is raised if an invalid process group is requested with pytest.raises(ValueError, match=r"Invalid process groups requested"): diff --git a/tests/unit_tests/test_static_benchmark.py b/tests/unit_tests/test_static_benchmark.py new file mode 100644 index 00000000000..a0f3f0fc7e7 --- /dev/null +++ b/tests/unit_tests/test_static_benchmark.py @@ -0,0 +1,106 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import argparse +import sys +from unittest import mock + +import pytest + +from tests.performance_tests.client import static_benchmark + + +@pytest.mark.parametrize( + ("batch_size", "data_parallel_size", "expected"), + [(1, 8, 8), (8, 8, 8), (32, 8, 32), (1, 4, 4), (128, 4, 128), (1, 1, 1)], +) +def test_get_warmup_batch_size(batch_size, data_parallel_size, expected): + assert static_benchmark._get_warmup_batch_size(batch_size, data_parallel_size) == expected + + +def test_parse_args_preserves_single_worker_default(monkeypatch): + monkeypatch.setattr(sys, "argv", ["static_benchmark.py"]) + + assert static_benchmark.parse_args().data_parallel_size == 1 + + +@pytest.mark.asyncio +async def test_run_batch_request_count_override(monkeypatch): + single_request = mock.AsyncMock(return_value=(512, 128, 0.1)) + monkeypatch.setattr(static_benchmark, "_single_request", single_request) + args = argparse.Namespace( + batch_size=1, model="gpt_583m", num_output_tokens=128, temperature=0.0 + ) + + inputs, outputs, latencies, _ = await static_benchmark._run_batch( + mock.sentinel.session, + args, + "http://localhost:5000/v1/completions", + ["prompt 0", "prompt 1"], + iter_start_index=1, + request_count=8, + ) + + assert single_request.await_count == 8 + assert [call.args[3] for call in single_request.await_args_list] == [ + "prompt 1", + "prompt 0", + "prompt 1", + "prompt 0", + "prompt 1", + "prompt 0", + "prompt 1", + "prompt 0", + ] + assert inputs == [512] * 8 + assert outputs == [128] * 8 + assert latencies == [0.1] * 8 + + single_request.reset_mock() + await static_benchmark._run_batch( + mock.sentinel.session, + args, + "http://localhost:5000/v1/completions", + ["prompt"], + iter_start_index=0, + ) + single_request.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_main_widens_only_warmup_batches_and_preserves_timed_prompts(monkeypatch): + calls = [] + + async def fake_run_batch(session, args, url, prompts, iter_start_index, request_count=None): + count = args.batch_size if request_count is None else request_count + calls.append((iter_start_index, count)) + return [512] * count, [128] * count, [0.1] * count, 1.0 + + class FakeClientSession: + async def __aenter__(self): + return mock.sentinel.session + + async def __aexit__(self, exc_type, exc_value, traceback): + return False + + monkeypatch.setattr(static_benchmark, "_run_batch", fake_run_batch) + monkeypatch.setattr(static_benchmark.aiohttp, "TCPConnector", mock.Mock()) + monkeypatch.setattr( + static_benchmark.aiohttp, "ClientSession", mock.Mock(return_value=FakeClientSession()) + ) + args = argparse.Namespace( + server_url="http://localhost:5000/v1", + model="gpt_583m", + batch_size=1, + data_parallel_size=8, + dataset="synthetic", + num_input_tokens=512, + num_output_tokens=128, + temperature=0.0, + num_warmup_iters=2, + num_iters=2, + ) + + summary = await static_benchmark.main(args) + + assert calls == [(0, 8), (8, 8), (2, 1), (3, 1)] + assert summary["batch_size"] == 1 diff --git a/tests/unit_tests/tokenizers/test_text_parsers.py b/tests/unit_tests/tokenizers/test_text_parsers.py new file mode 100644 index 00000000000..62ed6797149 --- /dev/null +++ b/tests/unit_tests/tokenizers/test_text_parsers.py @@ -0,0 +1,98 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +"""Parity tests for the ``/`` reasoning parsers. + +Ground truth for `NemotronV3ReasoningParser` is derived from vLLM's actual +implementation: + +- Base extraction: `BaseThinkingReasoningParser.extract_reasoning` in + `vllm/reasoning/basic_parsers.py` (used unmodified by `DeepSeekR1ReasoningParser` + for non-streaming extraction). Notably `final_content = content or None`, so an + empty string after a closing `` collapses to `None`, same as a missing + closing tag entirely. +- Override: `SuperV3ReasoningParser`/`UltraV3ReasoningParser.extract_reasoning` in + `super_v3_reasoning_parser.py`/`ultra_v3_reasoning_parser.py` (from + huggingface.co/nvidia/NVIDIA-Nemotron-3-{Super,Ultra}-*), which swaps all text + into content when `final_content is None` and either `enable_thinking is False` + or `force_nonempty_content is True`. + +""" + +import pytest + +from megatron.core.tokenizers.text.parsers import PARSER_MAPPING +from megatron.core.tokenizers.text.parsers.deepseek_r1_reasoning_parser import ( + DeepSeekR1ReasoningParser, +) +from megatron.core.tokenizers.text.parsers.nemotron_v3_reasoning_parser import ( + NemotronV3ReasoningParser, +) + +# (text, kwargs, expected_content, expected_info) +# `kwargs` is expanded into `parse(text, **kwargs)`; the override flags reach the +# parser inside `chat_template_kwargs`, exactly as the chat-completions endpoint +# forwards them from the request. +NEMOTRON_V3_CASES = [ + # No chat_template_kwargs override: behaves exactly like DeepSeekR1ReasoningParser. + ("hello", {}, "", {"reasoning": "hello"}), + ("helloworld", {}, "world", {"reasoning": "hello"}), + # Closing tag present but nothing follows it: vLLM's `content or None` treats + # this the same as a missing closing tag, so it is empty here too. + ("hello", {}, "", {"reasoning": "hello"}), + # No `` tag at all: vLLM assumes the whole string is reasoning. + ("just an answer", {}, "", {"reasoning": "just an answer"}), + # enable_thinking=False surfaces would-be-empty content as the reasoning text, + # for both the "unterminated" and "closes with nothing following" cases. + ("hello", {"chat_template_kwargs": {"enable_thinking": False}}, "hello", {}), + ("hello", {"chat_template_kwargs": {"enable_thinking": False}}, "hello", {}), + # force_nonempty_content=True has the same effect as enable_thinking=False. + ( + "hello", + {"chat_template_kwargs": {"force_nonempty_content": True}}, + "hello", + {}, + ), + ("hello", {"chat_template_kwargs": {"force_nonempty_content": True}}, "hello", {}), + # The override only fires when there would otherwise be no content. + ( + "helloworld", + {"chat_template_kwargs": {"enable_thinking": False}}, + "world", + {"reasoning": "hello"}, + ), + # Text preceding `` is discarded, override still applies past it. + ( + "prefixhello", + {"chat_template_kwargs": {"enable_thinking": False}}, + "hello", + {}, + ), + # enable_thinking=True (or omitted) must not trigger the override. + ( + "hello", + {"chat_template_kwargs": {"enable_thinking": True}}, + "", + {"reasoning": "hello"}, + ), +] + + +@pytest.mark.parametrize("text,kwargs,expected_content,expected_info", NEMOTRON_V3_CASES) +def test_nemotron_v3_reasoning_parser_matches_vllm(text, kwargs, expected_content, expected_info): + content, info = NemotronV3ReasoningParser.parse(text, **kwargs) + assert content == expected_content + assert info == expected_info + + +@pytest.mark.parametrize( + "text", ["hello", "helloworld", "hello", "just an answer"] +) +def test_nemotron_v3_reasoning_parser_without_override_matches_deepseek_r1(text): + """With no `enable_thinking`/`force_nonempty_content` kwargs, the Nemotron 3 + parser must be observably identical to the DeepSeek R1 parser it extends.""" + assert NemotronV3ReasoningParser.parse(text) == DeepSeekR1ReasoningParser.parse(text) + + +def test_parser_mapping_registers_nemotron_v3_reasoning(): + """Super and Ultra share identical reasoning-extraction logic upstream, so + both models are served by a single consolidated parser and registry key.""" + assert PARSER_MAPPING["nemotron-v3-reasoning"] is NemotronV3ReasoningParser diff --git a/tests/unit_tests/training/config/test_container_base.py b/tests/unit_tests/training/config/test_container_base.py index ea03a356cff..b6641665a63 100644 --- a/tests/unit_tests/training/config/test_container_base.py +++ b/tests/unit_tests/training/config/test_container_base.py @@ -203,7 +203,7 @@ def test_from_yaml_file_not_found(self): with pytest.raises(FileNotFoundError, match="YAML file not found"): TestConfigContainer.from_yaml("non_existent_file.yaml") - @patch("megatron.training.config.container.MultiStorageClientFeature.is_enabled") + @patch("megatron.core.msc_utils.MultiStorageClientFeature.is_enabled") @patch("omegaconf.OmegaConf") @patch("builtins.open", new_callable=mock_open) @patch("os.path.exists") @@ -243,7 +243,7 @@ def test_from_yaml_success(self, mock_exists, mock_file, mock_omegaconf, mock_ms assert result.name == "yaml_config" assert result.value == 500 - @patch("megatron.training.config.container.MultiStorageClientFeature.is_enabled") + @patch("megatron.core.msc_utils.MultiStorageClientFeature.is_enabled") @patch("os.path.exists") def test_from_yaml_with_mode(self, mock_exists, mock_msc): """Test from_yaml with different instantiation modes.""" diff --git a/tests/unit_tests/training/models/test_dist_utils.py b/tests/unit_tests/training/models/test_dist_utils.py index bfe875adcd5..353bc60ae4e 100644 --- a/tests/unit_tests/training/models/test_dist_utils.py +++ b/tests/unit_tests/training/models/test_dist_utils.py @@ -33,6 +33,9 @@ def _make_pg(): pg.pp.size.return_value = 1 pg.dp_cp.size.return_value = 1 pg.expt_dp.size.return_value = 1 + # With a single optimizer instance the intra-instance groups are the full groups. + pg.intra_dp_cp.size.return_value = 1 + pg.intra_expt_dp.size.return_value = 1 return pg @@ -500,6 +503,54 @@ def test_returns_list_of_wrapped_modules( assert isinstance(result, list) assert len(result) == 2 + @patch("megatron.training.models.dist_utils.DistributedDataParallel") + @patch("megatron.training.models.dist_utils.get_model_config") + @patch("torch.cuda.stream", new_callable=MagicMock) + @patch("torch.cuda.current_stream") + @patch("megatron.training.models.dist_utils.get_shared_capture_stream") + def test_uses_full_iteration_capture_stream_for_ddp_initialization( + self, mock_shared_stream, mock_current_stream, mock_stream_context, mock_config, mock_ddp + ): + mock_stream_context.return_value.__enter__ = Mock(return_value=None) + mock_stream_context.return_value.__exit__ = Mock(return_value=False) + shared_stream = mock_shared_stream.return_value + mock_config.return_value.cuda_graph_impl = "full_iteration" + + _ddp_wrap(self.model, False, self.ddp_config, False, pg_collection=self.pg) + + shared_stream.wait_stream.assert_called_once_with(mock_current_stream.return_value) + mock_stream_context.assert_called_once_with(shared_stream) + mock_current_stream.return_value.wait_stream.assert_called_once_with(shared_stream) + + @pytest.mark.parametrize("cuda_graph_impl", ["none", "local", "transformer_engine"]) + @patch("megatron.training.models.dist_utils.DistributedDataParallel") + @patch("megatron.training.models.dist_utils.get_model_config") + @patch("torch.cuda.stream", new_callable=MagicMock) + @patch("torch.cuda.current_stream") + @patch("torch.cuda.Stream") + @patch("megatron.training.models.dist_utils.get_shared_capture_stream") + def test_uses_dedicated_stream_for_other_cuda_graph_implementations( + self, + mock_shared_stream, + mock_stream, + mock_current_stream, + mock_stream_context, + mock_config, + mock_ddp, + cuda_graph_impl, + ): + mock_stream_context.return_value.__enter__ = Mock(return_value=None) + mock_stream_context.return_value.__exit__ = Mock(return_value=False) + mock_config.return_value.cuda_graph_impl = cuda_graph_impl + dedicated_stream = mock_stream.return_value + + _ddp_wrap(self.model, False, self.ddp_config, False, pg_collection=self.pg) + + mock_shared_stream.assert_not_called() + dedicated_stream.wait_stream.assert_called_once_with(mock_current_stream.return_value) + mock_stream_context.assert_called_once_with(dedicated_stream) + mock_current_stream.return_value.wait_stream.assert_called_once_with(dedicated_stream) + @patch("megatron.training.models.dist_utils.TorchFullyShardedDataParallel") @patch("megatron.training.models.dist_utils.HAVE_FSDP2", False) @patch("megatron.training.models.dist_utils.get_model_config") @@ -617,6 +668,9 @@ def setup_method(self): self.pg = _make_pg() self.pg.dp_cp.size.return_value = 4 self.pg.expt_dp.size.return_value = 2 + # Single optimizer instance, so the intra-instance groups match the full groups. + self.pg.intra_dp_cp.size.return_value = 4 + self.pg.intra_expt_dp.size.return_value = 2 self._opt_patcher = patch("megatron.training.models.dist_utils.DistributedOptimizer") self._opt = self._opt_patcher.start() self._opt.compute_full_param_layout.return_value = "LAYOUT" diff --git a/tests/unit_tests/training/test_freeze_all_layers.py b/tests/unit_tests/training/test_freeze_all_layers.py new file mode 100644 index 00000000000..fef00b0254c --- /dev/null +++ b/tests/unit_tests/training/test_freeze_all_layers.py @@ -0,0 +1,205 @@ +# Copyright (c) 2024-2026, NVIDIA CORPORATION. All rights reserved. + +"""Unit tests for the --freeze-all-layers helpers in megatron.training.training. + +These exercise ``_freeze_all_model_chunks`` and ``_forward_backward_grad_context`` +in isolation. Both operate on plain python objects (``requires_grad_``, an +attribute flip, and a grad context), so they run on CPU and need neither CUDA nor +a real Megatron model. The grad-context tests reproduce the PP>1 case that +motivates the fix: a recv_prev input activation with ``requires_grad=True``. +""" + +from contextlib import nullcontext +from types import SimpleNamespace + +import torch + +from megatron.training.training import _forward_backward_grad_context, _freeze_all_model_chunks + + +class _FakeRouter(torch.nn.Module): + """Stand-in for an MoE router that carries ``frozen_expert_bias`` (see + ``megatron/core/transformer/moe/router.py``).""" + + def __init__(self): + super().__init__() + self.gate = torch.nn.Linear(4, 2) + self.frozen_expert_bias = False + + +class _FakeModelChunk(torch.nn.Module): + """Minimal module tree: some trainable params plus a router submodule.""" + + def __init__(self): + super().__init__() + self.embedding = torch.nn.Linear(4, 8) + self.router = _FakeRouter() + self.output_layer = torch.nn.Linear(8, 4) + + +def _all_require_grad(module, value): + return all(p.requires_grad is value for p in module.parameters()) + + +def test_freezes_every_parameter(): + """All parameters across all chunks end up with requires_grad=False.""" + chunks = [_FakeModelChunk(), _FakeModelChunk()] + assert all(_all_require_grad(c, True) for c in chunks), "params start trainable" + + _freeze_all_model_chunks(chunks) + + assert all(_all_require_grad(c, False) for c in chunks) + + +def test_sets_frozen_expert_bias(): + """Modules exposing ``frozen_expert_bias`` are flipped to True; others are + left alone.""" + chunk = _FakeModelChunk() + assert chunk.router.frozen_expert_bias is False + + _freeze_all_model_chunks([chunk]) + + assert chunk.router.frozen_expert_bias is True + # A module without the attribute must not gain one. + assert not hasattr(chunk.embedding, "frozen_expert_bias") + + +def test_returns_same_list_object(): + """The helper freezes in place and returns the list it was given.""" + chunks = [_FakeModelChunk()] + + returned = _freeze_all_model_chunks(chunks) + + assert returned is chunks + + +def test_handles_multiple_routers_and_pp_style_chunks(): + """A VPP/PP-style list with several chunks, each with its own router, is + fully handled.""" + chunks = [_FakeModelChunk() for _ in range(3)] + + _freeze_all_model_chunks(chunks) + + for chunk in chunks: + assert _all_require_grad(chunk, False) + assert chunk.router.frozen_expert_bias is True + + +def test_empty_list_is_noop(): + """An empty chunk list is accepted and returned unchanged.""" + assert _freeze_all_model_chunks([]) == [] + + +def test_idempotent(): + """Applying the freeze twice keeps everything frozen.""" + chunk = _FakeModelChunk() + + _freeze_all_model_chunks([chunk]) + _freeze_all_model_chunks([chunk]) + + assert _all_require_grad(chunk, False) + assert chunk.router.frozen_expert_bias is True + + +# --------------------------------------------------------------------------- +# _forward_backward_grad_context +# +# The helper returns a ``(grad_context, forward_only)`` tuple that a frozen train +# step uses to mirror the eval forward pass: +# * grad_context is ``torch.no_grad()`` when frozen, else a no-op context; +# * forward_only is True when frozen, so the schedule skips the backward and +# finalize-grads collectives entirely. +# +# Why the grad_context matters for PP>1: on a non-first pipeline stage the input +# activation is received via ``create_tensor_recv_prev()`` +# (``megatron/core/pipeline_parallel/p2p_communication.py``), which allocates it +# with ``requires_grad=True`` so gradients can flow back to the prior stage. +# That means a fully *frozen* model still builds an autograd graph during forward +# -- purely because its input requires grad -- retaining activations for a +# backward that is never useful. ``forward_only`` alone does not prevent that +# graph (it only skips the backward call); ``torch.no_grad()`` is what suppresses +# it on frozen (e.g. teacher logits dump) runs. +# --------------------------------------------------------------------------- + + +def _recv_prev_activation(): + """A stand-in for the PP>1 stage input from ``create_tensor_recv_prev()``: + an activation tensor allocated with ``requires_grad=True``.""" + return torch.ones(2, 4, requires_grad=True) + + +def _forward_through_frozen_stage(chunk, recv_prev): + """Run a recv_prev activation through a frozen model chunk (Linear stack).""" + return chunk.output_layer(chunk.embedding(recv_prev)) + + +def test_frozen_returns_no_grad_and_forward_only(): + """With --freeze-all-layers: no_grad context and forward_only=True.""" + args = SimpleNamespace(freeze_all_layers=True) + + grad_context, forward_only = _forward_backward_grad_context(args) + + assert isinstance(grad_context, torch.no_grad) + assert forward_only is True + + +def test_unfrozen_returns_nullcontext_and_not_forward_only(): + """Without --freeze-all-layers: no-op context and forward_only=False.""" + args = SimpleNamespace(freeze_all_layers=False) + + grad_context, forward_only = _forward_backward_grad_context(args) + + assert isinstance(grad_context, nullcontext) + assert forward_only is False + + +def test_missing_flag_defaults_to_not_frozen(): + """A minimal args mock (no freeze_all_layers attr) is treated as unfrozen.""" + grad_context, forward_only = _forward_backward_grad_context(SimpleNamespace()) + + assert isinstance(grad_context, nullcontext) + assert forward_only is False + + +def test_pp_gt1_frozen_forward_without_context_still_builds_graph(): + """Regression: a frozen PP>1 stage builds a graph anyway, because the + recv_prev input requires grad. This is the situation the fix addresses.""" + chunk = _FakeModelChunk() + _freeze_all_model_chunks([chunk]) + assert _all_require_grad(chunk, False) # every parameter is frozen + + out = _forward_through_frozen_stage(chunk, _recv_prev_activation()) + + # Graph built despite all params frozen -- solely due to the recv_prev input. + assert out.requires_grad + assert out.grad_fn is not None + + +def test_pp_gt1_frozen_forward_under_context_skips_graph(): + """The fix: running the same frozen PP>1 forward under the freeze context + suppresses the autograd graph even though recv_prev requires grad.""" + chunk = _FakeModelChunk() + _freeze_all_model_chunks([chunk]) + args = SimpleNamespace(freeze_all_layers=True) + + grad_context, _ = _forward_backward_grad_context(args) + with grad_context: + assert not torch.is_grad_enabled() + out = _forward_through_frozen_stage(chunk, _recv_prev_activation()) + + assert out.grad_fn is None # no graph, no retained activations + assert torch.is_grad_enabled() # grad state restored on exit + + +def test_unfrozen_forward_builds_graph_normally(): + """The no-op context leaves normal (trainable) training untouched: a forward + over a recv_prev input still builds its graph.""" + chunk = _FakeModelChunk() # params trainable + args = SimpleNamespace(freeze_all_layers=False) + + grad_context, _ = _forward_backward_grad_context(args) + with grad_context: + assert torch.is_grad_enabled() + out = _forward_through_frozen_stage(chunk, _recv_prev_activation()) + + assert out.grad_fn is not None diff --git a/tests/unit_tests/training/test_param_norm.py b/tests/unit_tests/training/test_param_norm.py new file mode 100644 index 00000000000..27193ebf827 --- /dev/null +++ b/tests/unit_tests/training/test_param_norm.py @@ -0,0 +1,297 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import math +from types import SimpleNamespace + +import pytest +import torch + +from megatron.core.distributed import DistributedDataParallel, DistributedDataParallelConfig +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_with_transformer_engine_spec +from megatron.core.models.gpt.gpt_model import GPTModel +from megatron.core.optimizer import OptimizerConfig, get_megatron_optimizer +from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.training.utils import common_utils +from tests.unit_tests.test_utilities import Utils + + +def _build_tiny_moe_gpt( + tensor_parallel_size: int, + expert_parallel_size: int, + expert_tensor_parallel_size: int, + bf16: bool = False, + add_bias_linear: bool = False, +) -> GPTModel: + config = TransformerConfig( + num_layers=1, + hidden_size=8, + num_attention_heads=4, + ffn_hidden_size=16, + num_moe_experts=2, + moe_ffn_hidden_size=16, + # Shared experts do not support linear biases. + moe_shared_expert_intermediate_size=None if add_bias_linear else 16, + moe_router_topk=1, + moe_router_pre_softmax=True, + tensor_model_parallel_size=tensor_parallel_size, + expert_model_parallel_size=expert_parallel_size, + expert_tensor_parallel_size=expert_tensor_parallel_size, + sequence_parallel=tensor_parallel_size > 1, + use_cpu_initialization=True, + add_bias_linear=add_bias_linear, + normalization="RMSNorm", + moe_grouped_gemm=True, + bf16=bf16, + params_dtype=torch.bfloat16 if bf16 else torch.float32, + ) + model = GPTModel( + config=config, + transformer_layer_spec=get_gpt_layer_with_transformer_engine_spec( + num_experts=config.num_moe_experts, moe_grouped_gemm=True + ), + vocab_size=16, + max_sequence_length=8, + position_embedding_type="rope", + ) + if not add_bias_linear: + assert any(".shared_experts." in name for name, _ in model.named_parameters()) + return model.cuda() + + +def _fill_parameters_with_ones(model: GPTModel) -> None: + with torch.no_grad(): + for param in model.parameters(): + param.fill_(1.0) + + +@pytest.mark.parametrize( + ("tensor_parallel_size", "expert_parallel_size", "expert_tensor_parallel_size"), + ((2, 2, 1), (2, 1, 2), (4, 1, 2), (2, 1, 4)), + ids=("expert-parallel", "expert-tensor-parallel", "tp-larger-than-etp", "etp-larger-than-tp"), +) +def test_moe_param_norm_counts_each_logical_parameter_once( + monkeypatch, + tensor_parallel_size: int, + expert_parallel_size: int, + expert_tensor_parallel_size: int, +): + """Parameter norm should be invariant to expert and expert-tensor parallelism.""" + if Utils.world_size < 4 or Utils.world_size % 4 != 0: + pytest.skip("test requires a world size divisible by four") + + monkeypatch.setattr( + common_utils, "get_args", lambda: SimpleNamespace(use_megatron_fsdp=False, bf16=False) + ) + + try: + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, + expert_model_parallel_size=1, + expert_tensor_parallel_size=1, + ) + reference_model = _build_tiny_moe_gpt( + tensor_parallel_size=1, expert_parallel_size=1, expert_tensor_parallel_size=1 + ) + _fill_parameters_with_ones(reference_model) + expected_numel = sum(param.numel() for param in reference_model.parameters()) + expected_norm = math.sqrt(expected_numel) + reference_norm = common_utils.calc_params_l2_norm(reference_model) + + assert reference_norm == pytest.approx(expected_norm) + del reference_model + + Utils.initialize_model_parallel( + tensor_model_parallel_size=tensor_parallel_size, + expert_model_parallel_size=expert_parallel_size, + expert_tensor_parallel_size=expert_tensor_parallel_size, + ) + distributed_model = _build_tiny_moe_gpt( + tensor_parallel_size=tensor_parallel_size, + expert_parallel_size=expert_parallel_size, + expert_tensor_parallel_size=expert_tensor_parallel_size, + ) + _fill_parameters_with_ones(distributed_model) + + actual_norm = common_utils.calc_params_l2_norm(distributed_model) + + assert actual_norm == pytest.approx(expected_norm) + finally: + Utils.destroy_model_parallel() + + +@pytest.mark.parametrize("use_distributed_optimizer", (False, True), ids=("optimizer", "distopt")) +@pytest.mark.parametrize( + ("tensor_parallel_size", "expert_parallel_size", "expert_tensor_parallel_size"), + ((2, 2, 1), (2, 1, 2), (4, 1, 2), (2, 1, 4)), + ids=("expert-parallel", "expert-tensor-parallel", "tp-larger-than-etp", "etp-larger-than-tp"), +) +def test_moe_gradient_stats_and_clipping_count_each_logical_gradient_once( + tensor_parallel_size: int, + expert_parallel_size: int, + expert_tensor_parallel_size: int, + use_distributed_optimizer: bool, +): + """Gradient norm, clipping, and zero count should include each logical gradient once.""" + if Utils.world_size < 4 or Utils.world_size % 4 != 0: + pytest.skip("test requires a world size divisible by four") + + try: + Utils.initialize_model_parallel( + tensor_model_parallel_size=1, + expert_model_parallel_size=1, + expert_tensor_parallel_size=1, + ) + reference_model = _build_tiny_moe_gpt( + tensor_parallel_size=1, expert_parallel_size=1, expert_tensor_parallel_size=1, bf16=True + ) + expected_numel = sum(param.numel() for param in reference_model.parameters()) + expected_norm = math.sqrt(expected_numel) + del reference_model + + Utils.initialize_model_parallel( + tensor_model_parallel_size=tensor_parallel_size, + expert_model_parallel_size=expert_parallel_size, + expert_tensor_parallel_size=expert_tensor_parallel_size, + ) + model = _build_tiny_moe_gpt( + tensor_parallel_size=tensor_parallel_size, + expert_parallel_size=expert_parallel_size, + expert_tensor_parallel_size=expert_tensor_parallel_size, + bf16=True, + ) + ddp_config = DistributedDataParallelConfig( + grad_reduce_in_fp32=True, use_distributed_optimizer=use_distributed_optimizer + ) + model = DistributedDataParallel(model.config, ddp_config, model) + + max_norm = expected_norm / 2.0 + optimizer = get_megatron_optimizer( + OptimizerConfig( + optimizer="adam", + lr=0.0, + bf16=True, + clip_grad=max_norm, + log_num_zeros_in_grad=True, + use_distributed_optimizer=use_distributed_optimizer, + ), + [model], + ) + + for param in model.parameters(): + assert hasattr(param, "main_grad") + param.main_grad.zero_() + + found_inf = optimizer.prepare_grads() + assert not found_inf + assert optimizer.count_zeros() == expected_numel + + for param in model.parameters(): + param.main_grad.fill_(1.0) + + update_successful, actual_norm, actual_num_zeros = optimizer.step() + + assert update_successful + assert actual_num_zeros == 0 + actual_norm_value = ( + actual_norm.item() if isinstance(actual_norm, torch.Tensor) else actual_norm + ) + assert actual_norm_value == pytest.approx(expected_norm) + + expected_clip_coefficient = max_norm / (expected_norm + 1.0e-6) + grads_checked = 0 + for param in optimizer.get_parameters(): + if param.grad is None: + continue + torch.testing.assert_close( + param.grad, + torch.full_like(param.grad, expected_clip_coefficient), + rtol=1.0e-5, + atol=1.0e-6, + ) + grads_checked += 1 + assert grads_checked > 0 + finally: + Utils.destroy_model_parallel() + + +def test_layer_wise_muon_grad_norm_uses_expert_tp_group_for_row_parallel_bias(): + """LayerWise Muon must deduplicate replicated expert FC2 bias grads over ETP. + + With TP=2, EP=2, and ETP=1, every rank is ETP rank zero. The two EP ranks own + distinct row-parallel expert biases, so both gradients must contribute to the global + norm. Falling back to the regular TP rank drops the expert on TP rank one and + undercounts the squared norm by a factor of two. + """ + from megatron.core.optimizer.layer_wise_optimizer import LayerWiseDistributedOptimizer + from megatron.core.process_groups_config import ProcessGroupCollection + + if Utils.world_size < 4 or Utils.world_size % 4 != 0: + pytest.skip("test requires a world size divisible by four") + + tensor_parallel_size = 2 + expert_parallel_size = 2 + expert_tensor_parallel_size = 1 + + try: + Utils.initialize_model_parallel( + tensor_model_parallel_size=tensor_parallel_size, + expert_model_parallel_size=expert_parallel_size, + expert_tensor_parallel_size=expert_tensor_parallel_size, + ) + model = _build_tiny_moe_gpt( + tensor_parallel_size=tensor_parallel_size, + expert_parallel_size=expert_parallel_size, + expert_tensor_parallel_size=expert_tensor_parallel_size, + bf16=True, + add_bias_linear=True, + ) + + expert_fc2_biases = [ + param + for name, param in model.named_parameters() + if ".experts." in name and ".linear_fc2.bias" in name + ] + assert len(expert_fc2_biases) == model.config.num_moe_experts // expert_parallel_size + for parameter in expert_fc2_biases: + assert parameter.ndim == 1 + assert parameter.allreduce is False + assert parameter.tensor_model_parallel is False + + model = DistributedDataParallel( + model.config, DistributedDataParallelConfig(use_distributed_optimizer=False), model + ) + pg_collection = ProcessGroupCollection.use_mpu_process_groups() + optimizer = get_megatron_optimizer( + OptimizerConfig( + optimizer="muon", + lr=0.0, + weight_decay=0.0, + bf16=True, + use_distributed_optimizer=False, + use_layer_wise_distributed_optimizer=True, + muon_tp_mode="duplicated", + ), + [model], + use_gloo_process_groups=False, + pg_collection=pg_collection, + ) + + assert isinstance(optimizer, LayerWiseDistributedOptimizer) + assert pg_collection.tp.size() == tensor_parallel_size + assert pg_collection.expt_tp.size() == expert_tensor_parallel_size + + for parameter in model.parameters(): + parameter.main_grad.zero_() + for parameter in expert_fc2_biases: + parameter.main_grad.fill_(1.0) + assert optimizer.prepare_grads() is False + + actual_norm = optimizer.get_grad_norm() + actual_norm_value = ( + actual_norm.item() if isinstance(actual_norm, torch.Tensor) else actual_norm + ) + expected_norm = math.sqrt(model.config.num_moe_experts * model.config.hidden_size) + + assert actual_norm_value == pytest.approx(expected_norm) + finally: + Utils.destroy_model_parallel() diff --git a/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py b/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py index 8f0f039d031..bdab0c9c845 100644 --- a/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py +++ b/tests/unit_tests/transformer/experimental_attention_variant/test_absorbed_mla.py @@ -129,6 +129,7 @@ def get_mock_mla_config( """Create test config with all attributes used in MLA.""" return MLATransformerConfig( multi_latent_attention=True, + qk_layernorm=True, hidden_size=7168, num_attention_heads=128, q_lora_rank=1536, diff --git a/tests/unit_tests/transformer/experimental_attention_variant/test_dsa_tilelang_kernels.py b/tests/unit_tests/transformer/experimental_attention_variant/test_dsa_tilelang_kernels.py new file mode 100644 index 00000000000..2a9e151d4ee --- /dev/null +++ b/tests/unit_tests/transformer/experimental_attention_variant/test_dsa_tilelang_kernels.py @@ -0,0 +1,1284 @@ +# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved. + +import math +from types import SimpleNamespace + +import pytest +import torch + +from megatron.core.packed_seq_params import PackedSeqParams +from megatron.core.transformer.experimental_attention_variant import ( + dsa_indexer_loss, + dsa_masking, + dsa_tilelang_kernels, +) +from megatron.core.transformer.experimental_attention_variant.ops import ( + indexer, + sparse_mla, + tilelang_dsa, + tilelang_indexer_bwd, + tilelang_indexer_fwd, + tilelang_indexer_loss, + tilelang_sparse_mla_bwd, + tilelang_utils, +) + + +def test_run_fused_qk_topk_forwards_to_tilelang_backend(monkeypatch): + q = torch.empty(2, 1, 3, 4) + k = torch.empty(5, 1, 4) + weights = torch.empty(2, 1, 3) + starts = torch.tensor([0, 1], dtype=torch.int32) + ends = torch.tensor([3, 5], dtype=torch.int32) + expected_indices = torch.tensor([[[2, 1], [4, 3]]], dtype=torch.int32) + call = {} + + def fake_run_fused_qk_topk( + q_arg, k_arg, weights_arg, index_topk, starts_arg, ends_arg, block_size, use_relu, **kwargs + ): + call.update( + q=q_arg, + k=k_arg, + weights=weights_arg, + index_topk=index_topk, + starts=starts_arg, + ends=ends_arg, + block_size=block_size, + use_relu=use_relu, + kwargs=kwargs, + ) + return expected_indices + + monkeypatch.setattr( + dsa_tilelang_kernels.tilelang_dsa, "run_fused_qk_topk", fake_run_fused_qk_topk + ) + + result = dsa_tilelang_kernels.run_fused_qk_topk( + q, + k, + weights, + index_topk=2, + starts=starts, + ends=ends, + block_size=8, + use_relu=False, + use_local_indexer_varlen=True, + ) + + indices, topk_length = result + assert indices is expected_indices + assert topk_length is None + assert call["q"] is q + assert call["k"] is k + assert call["weights"] is weights + assert call["index_topk"] == 2 + assert call["starts"] is starts + assert call["ends"] is ends + assert call["block_size"] == 8 + assert call["use_relu"] is False + assert call["kwargs"]["use_local_indexer_varlen"] is True + assert call["kwargs"]["cp_size"] == 1 + + +def test_run_fused_qk_topk_preserves_unavailable_backend(monkeypatch): + def fake_run_fused_qk_topk(*_args, **_kwargs): + return None + + monkeypatch.setattr( + dsa_tilelang_kernels.tilelang_dsa, "run_fused_qk_topk", fake_run_fused_qk_topk + ) + + result = dsa_tilelang_kernels.run_fused_qk_topk( + torch.empty(2, 1, 3, 4), + torch.empty(5, 1, 4), + torch.empty(2, 1, 3), + index_topk=2, + starts=torch.tensor([0, 1], dtype=torch.int32), + ends=torch.tensor([3, 5], dtype=torch.int32), + block_size=8, + ) + + assert result is None + + +def test_tilelang_packed_cp_indexer_inputs_segment_keys_and_bounds(): + packed_seq_params = PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=torch.tensor([0, 8, 24], dtype=torch.int32), + cu_seqlens_kv=torch.tensor([0, 8, 24], dtype=torch.int32), + max_seqlen_q=16, + max_seqlen_kv=16, + ) + query_positions = torch.tensor([0, 1, 6, 7, 8, 9, 10, 11, 20, 21, 22, 23]) + starts = torch.tensor([0] * 4 + [8] * 8, dtype=torch.int32) + ends = (query_positions + 1).to(torch.int32) + index_k = torch.arange(24, dtype=torch.float32).view(24, 1) + + segmented_k, local_starts, local_ends, source_indices = ( + tilelang_dsa._build_packed_cp_indexer_inputs( + index_k, + starts, + ends, + packed_seq_params=packed_seq_params, + cp_size=2, + cp_rank=0, + single_packed_thd_sequence=False, + local_query_start=0, + local_query_len=12, + ) + ) + + expected_sources = torch.tensor( + [0, 1, *range(8), *range(8, 12), *range(8, 24)], dtype=torch.int64 + ) + torch.testing.assert_close(source_indices, expected_sources) + torch.testing.assert_close(segmented_k[:, 0], expected_sources.to(torch.float32)) + torch.testing.assert_close( + local_starts, torch.tensor([0, 0, 2, 2, 10, 10, 10, 10, 14, 14, 14, 14], dtype=torch.int32) + ) + torch.testing.assert_close( + local_ends, torch.tensor([1, 2, 9, 10, 11, 12, 13, 14, 27, 28, 29, 30], dtype=torch.int32) + ) + + +def test_tilelang_packed_cp_indexer_remaps_segmented_topk(monkeypatch): + packed_seq_params = PackedSeqParams( + qkv_format="thd", + cu_seqlens_q=torch.tensor([0, 8, 24], dtype=torch.int32), + cu_seqlens_kv=torch.tensor([0, 8, 24], dtype=torch.int32), + max_seqlen_q=16, + max_seqlen_kv=16, + ) + query_positions = torch.tensor([0, 1, 6, 7, 8, 9, 10, 11, 20, 21, 22, 23]) + starts = torch.tensor([0] * 4 + [8] * 8, dtype=torch.int32) + ends = (query_positions + 1).to(torch.int32) + seen = {} + + def fake_lighting_indexer_indices( + index_q, index_k, index_w, starts_arg, ends_arg, index_topk, use_relu=True + ): + del index_q, index_w, use_relu + seen["key"] = index_k[:, 0].clone() + seen["starts"] = starts_arg.clone() + seen["ends"] = ends_arg.clone() + offsets = torch.arange(index_topk, dtype=torch.int32).view(1, -1) + return ends_arg.view(-1, 1) - 1 - offsets + + monkeypatch.setattr(tilelang_dsa, "lighting_indexer_indices", fake_lighting_indexer_indices) + topk = tilelang_dsa.fused_qk_topk_lighting( + torch.ones((12, 1, 1, 1), dtype=torch.bfloat16), + torch.arange(24, dtype=torch.bfloat16).view(24, 1, 1), + torch.ones((12, 1, 1)), + index_topk=3, + starts=starts, + ends=ends, + block_size=12, + use_local_indexer_varlen=True, + packed_seq_params=packed_seq_params, + cp_size=2, + ) + + expected = [] + for position, sequence_start in zip(query_positions.tolist(), starts.tolist()): + row = list(range(position, max(sequence_start - 1, position - 3), -1)) + expected.append(row + [-1] * (3 - len(row))) + torch.testing.assert_close(topk, torch.tensor([expected], dtype=torch.int32)) + assert seen["key"].numel() == 30 + torch.testing.assert_close( + seen["starts"], + torch.tensor([0, 0, 2, 2, 10, 10, 10, 10, 14, 14, 14, 14], dtype=torch.int32), + ) + + +def test_run_fused_qk_topk_with_loss_preserves_unavailable_backend(monkeypatch): + def fake_run_fused_qk_topk_with_loss(**kwargs): + return None + + monkeypatch.setattr( + dsa_tilelang_kernels.tilelang_dsa, + "run_fused_qk_topk_with_loss", + fake_run_fused_qk_topk_with_loss, + ) + + result = dsa_tilelang_kernels.run_fused_qk_topk_with_loss( + q=torch.empty(2, 1, 3, 4), + k=torch.empty(5, 1, 4), + weights=torch.empty(2, 1, 3), + index_topk=2, + starts=torch.tensor([0, 1], dtype=torch.int32), + ends=torch.tensor([3, 5], dtype=torch.int32), + block_size=8, + query=torch.empty(2, 1, 3, 4), + key=torch.empty(5, 1, 1, 4), + softmax_scale=0.5, + loss_coeff=0.1, + pg_collection=SimpleNamespace(), + config=SimpleNamespace(), + use_local_indexer_varlen=True, + ) + + assert result is None + + +def test_run_fused_qk_topk_with_loss_adds_empty_topk_length(monkeypatch): + q = torch.empty(2, 1, 3, 4) + k = torch.empty(5, 1, 4) + weights = torch.empty(2, 1, 3) + starts = torch.tensor([0, 1], dtype=torch.int32) + ends = torch.tensor([3, 5], dtype=torch.int32) + query = torch.empty(2, 1, 3, 4) + key = torch.empty(5, 1, 1, 4) + query_valid_rows = torch.tensor([[True, False]]) + pg_collection = SimpleNamespace() + expected_indices = torch.tensor([[[2, 1], [4, 3]]], dtype=torch.int32) + expected_loss = torch.tensor(1.25) + call = {} + + def fake_run_fused_qk_topk_with_loss(**kwargs): + call.update(kwargs) + return expected_indices, expected_loss + + monkeypatch.setattr( + dsa_tilelang_kernels.tilelang_dsa, + "run_fused_qk_topk_with_loss", + fake_run_fused_qk_topk_with_loss, + ) + + result = dsa_tilelang_kernels.run_fused_qk_topk_with_loss( + q=q, + k=k, + weights=weights, + index_topk=2, + starts=starts, + ends=ends, + block_size=8, + query=query, + key=key, + softmax_scale=0.5, + loss_coeff=0.1, + pg_collection=pg_collection, + query_valid_rows=query_valid_rows, + calculate_per_token_loss=True, + use_relu=False, + config=SimpleNamespace(), + use_local_indexer_varlen=True, + ) + + indices, topk_length, indexer_loss = result + assert indices is expected_indices + assert topk_length is None + assert indexer_loss is expected_loss + assert call["q"] is q + assert call["k"] is k + assert call["weights"] is weights + assert call["index_topk"] == 2 + assert call["starts"] is starts + assert call["ends"] is ends + assert call["block_size"] == 8 + assert call["query"] is query + assert call["key"] is key + assert call["softmax_scale"] == 0.5 + assert call["loss_coeff"] == 0.1 + assert call["pg_collection"] is pg_collection + assert call["query_valid_rows"] is query_valid_rows + assert call["calculate_per_token_loss"] is True + assert call["use_relu"] is False + + +def test_run_fused_absorbed_sparse_attention_forwards_to_tilelang_backend(monkeypatch): + query = torch.empty(2, 1, 3, 4) + key = torch.empty(5, 1, 1, 4) + topk_indices = torch.tensor([[[0, 1], [1, 99]]], dtype=torch.int32) + topk_length = torch.tensor([[2, 1]], dtype=torch.int32) + expected_output = torch.empty(2, 1, 3, 4) + call = {} + + def fake_run_fused_absorbed_sparse_attention( + query_arg, key_arg, topk_indices_arg, softmax_scale, v_channels + ): + call.update( + query=query_arg, + key=key_arg, + topk_indices=topk_indices_arg, + softmax_scale=softmax_scale, + v_channels=v_channels, + ) + return expected_output + + monkeypatch.setattr( + dsa_tilelang_kernels.tilelang_dsa, + "run_fused_absorbed_sparse_attention", + fake_run_fused_absorbed_sparse_attention, + ) + + result = dsa_tilelang_kernels.run_fused_absorbed_sparse_attention( + query, key, topk_indices, softmax_scale=0.5, v_channels=4, topk_length=topk_length + ) + + assert result is expected_output + assert call["query"] is query + assert call["key"] is key + torch.testing.assert_close( + call["topk_indices"], torch.tensor([[[0, 1], [1, -1]]], dtype=torch.int32) + ) + assert call["softmax_scale"] == 0.5 + assert call["v_channels"] == 4 + + +def test_indexer_topk_helpers_mask_invalid_entries(): + logits = torch.tensor([[1.0, 3.0, float("-inf")], [0.0, 2.0, 1.0]]) + requested_indices = torch.tensor([[1, -1, 4], [0, 2, 1]], dtype=torch.int32) + + gathered = indexer.pytorch_extract_topk_scores(logits, requested_indices) + + assert torch.equal(gathered[0], torch.tensor([3.0, float("-inf"), float("-inf")])) + assert torch.equal(gathered[1], torch.tensor([0.0, 1.0, 2.0])) + + topk_scores, topk_indices = indexer._select_topk_from_logits(logits, topk=4) + assert topk_scores.shape == (2, 3) + assert topk_indices.shape == (2, 3) + assert topk_indices.dtype == torch.int32 + assert -1 in topk_indices[0].tolist() + + empty_scores, empty_indices = indexer._select_topk_from_logits(torch.empty(2, 0), topk=3) + assert empty_scores.shape == (2, 0) + assert empty_indices.shape == (2, 0) + assert empty_indices.dtype == torch.int32 + + +def test_sparse_mla_head_mask_helpers(): + indices = torch.tensor([[[0, -1], [-1, -1]], [[1, 2], [3, -1]]], dtype=torch.int32) + + valid_heads = sparse_mla._valid_head_mask(indices, num_heads=4) + + assert torch.equal( + valid_heads, torch.tensor([[True, True, False, False], [True, True, True, True]]) + ) + + tensor = torch.arange(2 * 4 * 3, dtype=torch.float32).view(2, 4, 3) + zeroed = sparse_mla._zero_invalid_heads(tensor, valid_heads) + + assert torch.equal(zeroed[0, :2], tensor[0, :2]) + assert torch.equal(zeroed[0, 2:], torch.zeros_like(tensor[0, 2:])) + assert torch.equal(zeroed[1], tensor[1]) + + batched_indices = indices.unsqueeze(0) + batched_valid_heads = sparse_mla._valid_head_mask(batched_indices, num_heads=4) + assert torch.equal(batched_valid_heads, valid_heads.unsqueeze(0)) + + batched_tensor = tensor.unsqueeze(0) + batched_zeroed = sparse_mla._zero_invalid_heads(batched_tensor, batched_valid_heads) + assert torch.equal(batched_zeroed, zeroed.unsqueeze(0)) + + with pytest.raises(RuntimeError, match="heads must be divisible"): + sparse_mla._valid_head_mask(indices, num_heads=3) + + +def test_tilelang_dsa_sanitize_helper(): + topk_indices = torch.tensor([[0, 2, 5], [-1, 3, 4]], dtype=torch.int32) + topk_scores = torch.tensor([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]) + starts = torch.tensor([1, 3], dtype=torch.int32) + ends = torch.tensor([5, 4], dtype=torch.int32) + + sanitized_indices, sanitized_scores = tilelang_dsa._sanitize_fused_topk_outputs( + topk_indices, starts, ends, topk_scores + ) + + assert torch.equal(sanitized_indices, torch.tensor([[-1, 2, -1], [-1, 3, -1]])) + assert torch.equal( + torch.isneginf(sanitized_scores), torch.tensor([[True, False, True], [True, False, True]]) + ) + + +def test_tilelang_dsa_scratch_cache_reuses_buffers(monkeypatch): + tilelang_dsa._DSA_SCRATCH_CACHE.clear() + monkeypatch.setattr(tilelang_dsa, "_DSA_SCRATCH_CACHE_TOTAL_BYTES", 0) + monkeypatch.setattr(tilelang_dsa, "_DSA_SCRATCH_CACHE_MAX_ENTRIES", 1) + monkeypatch.setattr(tilelang_dsa, "_DSA_SCRATCH_CACHE_MAX_BYTES", 1024) + + first = tilelang_dsa._get_scratch_buffer("a", (2,), torch.float32, torch.device("cpu")) + first.fill_(3.0) + reused = tilelang_dsa._get_scratch_buffer("a", (2,), torch.float32, torch.device("cpu")) + second = tilelang_dsa._get_scratch_buffer("b", (2,), torch.float32, torch.device("cpu")) + + assert reused is first + assert torch.equal(reused, torch.full((2,), 3.0)) + assert list(tilelang_dsa._DSA_SCRATCH_CACHE) == [ + ("b", (2,), torch.float32, torch.device("cpu")) + ] + assert tilelang_dsa._DSA_SCRATCH_CACHE_TOTAL_BYTES == second.numel() * second.element_size() + + +def test_tilelang_kernel_helper_caches_and_env_parsing(monkeypatch): + monkeypatch.delenv("MCORE_DSA_TILELANG_KERNEL_CACHE_MAX", raising=False) + assert tilelang_utils._env_int("MCORE_DSA_TILELANG_KERNEL_CACHE_MAX", 7) == 7 + + monkeypatch.setenv("MCORE_DSA_TILELANG_KERNEL_CACHE_MAX", "bad") + assert tilelang_utils._env_int("MCORE_DSA_TILELANG_KERNEL_CACHE_MAX", 7) == 7 + + monkeypatch.setenv("MCORE_DSA_TILELANG_KERNEL_CACHE_MAX", "-3") + assert tilelang_utils._env_int("MCORE_DSA_TILELANG_KERNEL_CACHE_MAX", 7) == 7 + + monkeypatch.setenv("MCORE_DSA_TILELANG_KERNEL_CACHE_MAX", "2") + assert tilelang_utils._env_int("MCORE_DSA_TILELANG_KERNEL_CACHE_MAX", 7) == 2 + + # Shared numeric/layout helpers now live in tilelang_utils. + assert tilelang_utils._next_power_of_two(0) == 1 + assert tilelang_utils._next_power_of_two(9) == 16 + assert tilelang_utils._round_up(9, 4) == 12 + assert tilelang_utils._round_up(9, 1) == 9 + assert tilelang_utils._normalize_sm_scale(None) is None + assert tilelang_utils._normalize_sm_scale(torch.tensor(0.5)) == 0.5 + + # Kernel-specific helpers stay with their modules. + assert tilelang_indexer_bwd._canonical_topk(33) == 64 + assert tilelang_indexer_bwd.is_supported_indexer_bwd_head_count(8) + assert tilelang_indexer_bwd.is_supported_indexer_bwd_head_count(64) + assert not tilelang_indexer_bwd.is_supported_indexer_bwd_head_count(7) + assert not tilelang_indexer_bwd.is_supported_indexer_bwd_head_count(72) + assert tilelang_sparse_mla_bwd._normalize_block_h(12) == 16 + assert tilelang_sparse_mla_bwd._normalize_block_h(40) == 32 + assert tilelang_sparse_mla_bwd._normalize_block_h(80) == 64 + assert tilelang_dsa._is_supported_sparse_mla_head_count(16) + assert tilelang_dsa._is_supported_sparse_mla_head_count(32) + assert tilelang_dsa._is_supported_sparse_mla_head_count(64) + assert tilelang_dsa._is_supported_sparse_mla_head_count(128) + assert tilelang_dsa._is_supported_sparse_mla_head_count(256, kv_group=2) + assert not tilelang_dsa._is_supported_sparse_mla_head_count(96) + assert not tilelang_dsa._is_supported_sparse_mla_head_count(96, kv_group=0) + # head_kv that is not a power of two >= 16 pads to a larger head dim in the kernels and + # would index past the real head count, so it must decline to the unfused path. + assert not tilelang_dsa._is_supported_sparse_mla_head_count(8) + assert not tilelang_dsa._is_supported_sparse_mla_head_count(48) + assert not tilelang_dsa._is_supported_sparse_mla_head_count(192) + assert not tilelang_dsa._is_supported_sparse_mla_head_count(96, kv_group=2) + + +def test_sparse_mla_canonicalizes_size_one_batch_stride_without_copy(): + tensor_sbhd = torch.empty(256, 1, 1, 4) + tensor_bshd = tensor_sbhd.permute(1, 0, 2, 3) + + assert tensor_bshd.is_contiguous() + assert tensor_bshd.stride(0) != tensor_bshd.numel() + + normalized = sparse_mla._canonicalize_batch_stride(tensor_bshd) + + assert normalized.stride(0) == normalized.numel() + assert normalized.data_ptr() == tensor_bshd.data_ptr() + + +def test_indexer_bwd_returns_grad_k_in_index_k_dtype(monkeypatch): + captured = {} + + def fake_kernel(index_q, index_k, weights, topk_indices, grad_scores, grad_q, grad_w, grad_k): + del index_q, index_k, weights, topk_indices, grad_scores + captured["grad_k_kernel_dtype"] = grad_k.dtype + grad_q.fill_(1) + grad_w.fill_(2) + grad_k.fill_(3) + + monkeypatch.setattr(tilelang_indexer_bwd, "require_tilelang", lambda: None) + monkeypatch.setattr( + tilelang_indexer_bwd, "_get_indexer_bwd_kernel", lambda *_args, **_kwargs: fake_kernel + ) + + index_q = torch.empty((2, 8, 4), dtype=torch.bfloat16) + index_k = torch.empty((3, 4), dtype=torch.bfloat16) + weights = torch.empty((2, 8), dtype=torch.float32) + topk_indices = torch.zeros((2, 1), dtype=torch.int32) + grad_scores = torch.empty((2, 1), dtype=torch.float32) + + _, _, grad_k = tilelang_indexer_bwd.indexer_bwd_interface( + index_q, weights, index_k, topk_indices, grad_scores + ) + + assert captured["grad_k_kernel_dtype"] == torch.float32 + assert grad_k.dtype == index_k.dtype + torch.testing.assert_close(grad_k.float(), torch.full_like(grad_k, 3, dtype=torch.float32)) + + +def test_sparse_mla_delta_pads_partial_sequence_tile(monkeypatch): + seq_len = 65 + padded_seq_len = 96 + heads = 2 + dim = 4 + o = torch.arange(seq_len * heads * dim, dtype=torch.float32).view(seq_len, heads, dim) + do = torch.full_like(o, 2.0) + + monkeypatch.setattr(tilelang_sparse_mla_bwd, "require_tilelang", lambda: None) + + def fake_get_preprocess_kernel(H, D): + assert H == heads + assert D == dim + + def fake_preprocess_kernel(o_arg, do_arg): + assert o_arg.shape == (1, padded_seq_len, heads, dim) + assert do_arg.shape == (1, padded_seq_len, heads, dim) + assert torch.equal(o_arg[:, :seq_len], o.unsqueeze(0)) + assert torch.equal(do_arg[:, :seq_len], do.unsqueeze(0)) + assert torch.equal(o_arg[:, seq_len:], torch.zeros_like(o_arg[:, seq_len:])) + assert torch.equal(do_arg[:, seq_len:], torch.zeros_like(do_arg[:, seq_len:])) + return torch.arange(padded_seq_len * heads, dtype=torch.float32).view( + 1, padded_seq_len, heads + ) + + return fake_preprocess_kernel + + monkeypatch.setattr( + tilelang_sparse_mla_bwd, "_get_preprocess_kernel", fake_get_preprocess_kernel + ) + + delta = tilelang_sparse_mla_bwd.sparse_mla_delta(o.contiguous(), do.contiguous()) + + assert delta.shape == (seq_len, heads) + assert delta.is_contiguous() + expected = torch.arange(padded_seq_len * heads, dtype=torch.float32).view( + padded_seq_len, heads + )[:seq_len] + torch.testing.assert_close(delta, expected) + + +def test_lighting_indexer_indices_preserves_single_head_weight_axis(monkeypatch): + seen = {} + + def fake_indexer_fwd_interface( + index_q, index_k, weights, cu_seqlen_ks, cu_seqlen_ke, clean_logits, use_relu + ): + del index_q, index_k, cu_seqlen_ks, cu_seqlen_ke + seen["weights_shape"] = weights.shape + seen["clean_logits"] = clean_logits + seen["use_relu"] = use_relu + return torch.arange(6, dtype=torch.float32).view(2, 3) + + monkeypatch.setattr(indexer, "indexer_fwd_interface", fake_indexer_fwd_interface) + + topk_indices = indexer.lighting_indexer_indices( + index_q=torch.empty(2, 1, 4), + index_k=torch.empty(3, 4), + weights=torch.ones(2, 1), + cu_seqlen_ks=torch.zeros(2, dtype=torch.int32), + cu_seqlen_ke=torch.full((2,), 3, dtype=torch.int32), + topk=2, + use_relu=False, + ) + + assert seen["weights_shape"] == (2, 1) + assert seen["clean_logits"] is True + assert seen["use_relu"] is False + torch.testing.assert_close(topk_indices, torch.tensor([[2, 1], [2, 1]], dtype=torch.int32)) + + +def _skip_if_real_tilelang_indexer_unavailable(): + if not torch.cuda.is_available(): + pytest.skip("CUDA is required for TileLang indexer parity tests") + if not indexer.HAVE_TILELANG_INDEXER: + pytest.skip("TileLang indexer forward/backward kernels are unavailable") + + +def _pytorch_indexer_scores(index_q, index_k, weights, *, use_relu): + per_head_scores = torch.einsum("qhd,kd->qkh", index_q.float(), index_k.float()) + if use_relu: + per_head_scores = per_head_scores.relu() + return (per_head_scores * weights.float().unsqueeze(1)).sum(dim=-1) + + +@pytest.mark.parametrize("use_relu", [False, True]) +def test_tilelang_indexer_forward_matches_pytorch(use_relu): + _skip_if_real_tilelang_indexer_unavailable() + torch.manual_seed(1234) + + device = torch.device("cuda") + q_len, k_len, heads, dim = 5, 19, 8, 16 + index_q = (torch.randn(q_len, heads, dim, device=device) * 0.25).to(torch.bfloat16) + index_k = (torch.randn(k_len, dim, device=device) * 0.25).to(torch.bfloat16) + weights = torch.randn(q_len, heads, dtype=torch.float32, device=device) * 0.25 + starts = torch.tensor([0, 1, 3, 5, 8], dtype=torch.int32, device=device) + ends = torch.tensor([7, 10, 13, 17, 19], dtype=torch.int32, device=device) + + actual = indexer.indexer_fwd_interface( + index_q, index_k, weights, starts, ends, clean_logits=True, use_relu=use_relu + ) + expected = _pytorch_indexer_scores(index_q, index_k, weights, use_relu=use_relu) + key_positions = torch.arange(k_len, device=device) + valid = (key_positions.unsqueeze(0) >= starts.unsqueeze(1)) & ( + key_positions.unsqueeze(0) < ends.unsqueeze(1) + ) + + torch.testing.assert_close(actual[valid], expected[valid], rtol=2e-2, atol=2e-2) + assert torch.isneginf(actual[~valid]).all() + + +@pytest.mark.parametrize("use_relu", [False, True]) +def test_tilelang_indexer_backward_matches_pytorch(use_relu): + _skip_if_real_tilelang_indexer_unavailable() + torch.manual_seed(5678) + + device = torch.device("cuda") + q_len, k_len, heads, dim = 4, 32, 8, 16 + index_q = (torch.randn(q_len, heads, dim, device=device) * 0.25).to(torch.bfloat16) + index_k = (torch.randn(k_len, dim, device=device) * 0.25).to(torch.bfloat16) + weights = torch.randn(q_len, heads, dtype=torch.float32, device=device) * 0.25 + topk_indices = torch.tensor( + [ + [0, 2, 4, 6, 8, 10, -1], + [1, 3, 5, 7, 9, 11, -1], + [12, 14, 16, 18, 20, 22, -1], + [13, 15, 17, 19, 21, 23, -1], + ], + dtype=torch.int32, + device=device, + ) + grad_scores = torch.randn(topk_indices.shape, dtype=torch.float32, device=device) + grad_scores.masked_fill_(topk_indices < 0, 0.0) + + actual_grad_q, actual_grad_w, actual_grad_k = indexer.indexer_bwd_interface( + index_q, weights, index_k, topk_indices, grad_scores, use_relu=use_relu + ) + + reference_q = index_q.detach().clone().requires_grad_(True) + reference_k = index_k.detach().clone().requires_grad_(True) + reference_w = weights.detach().clone().requires_grad_(True) + reference_scores = _pytorch_indexer_scores( + reference_q, reference_k, reference_w, use_relu=use_relu + ) + valid = topk_indices >= 0 + selected_scores = reference_scores.gather(1, topk_indices.clamp_min(0).long()) + selected_scores = selected_scores.masked_fill(~valid, 0.0) + (selected_scores * grad_scores).sum().backward() + + torch.testing.assert_close(actual_grad_q, reference_q.grad, rtol=5e-2, atol=5e-2) + torch.testing.assert_close(actual_grad_w, reference_w.grad, rtol=5e-2, atol=5e-2) + torch.testing.assert_close(actual_grad_k, reference_k.grad, rtol=5e-2, atol=5e-2) + + +def test_shared_topk_sort_uses_explicit_validity_mask(): + indices = torch.tensor([[5, 1, 7, 3]], dtype=torch.int32) + scores = torch.tensor([[0.5, 0.1, 0.7, 0.3]]) + valid = torch.tensor([[True, False, True, False]]) + + sorted_indices, sorted_scores = dsa_masking.sort_topk_by_index( + indices, valid, sk=8, topk_scores=scores + ) + + torch.testing.assert_close(sorted_indices, torch.tensor([[5, 7, -1, -1]], dtype=torch.int32)) + torch.testing.assert_close(sorted_scores[:, :2], torch.tensor([[0.5, 0.7]])) + assert torch.isneginf(sorted_scores[:, 2:]).all() + + +def test_tilelang_kernel_getters_reuse_cached_builders(monkeypatch): + def make_fake_kernel(): + return lambda *_args, **_kwargs: None + + try: + monkeypatch.setattr(tilelang_utils, "_TILELANG_KERNEL_CACHE_MAX", 1) + tilelang_indexer_fwd._tilelang_indexer_fwd_kernel_cache.clear() + tilelang_indexer_fwd._tilelang_indexer_clean_logits_kernel_cache.clear() + fwd_builds = [] + clean_builds = [] + + def fake_indexer_builder(**kwargs): + fwd_builds.append(kwargs) + return make_fake_kernel() + + def fake_clean_builder(**kwargs): + clean_builds.append(kwargs) + return make_fake_kernel() + + monkeypatch.setattr(tilelang_indexer_fwd, "tl_indexer_fwd_impl", fake_indexer_builder) + monkeypatch.setattr(tilelang_indexer_fwd, "clean_logits_", fake_clean_builder) + + first = tilelang_indexer_fwd._get_indexer_fwd_kernel(2, 4) + second = tilelang_indexer_fwd._get_indexer_fwd_kernel(2, 4) + third = tilelang_indexer_fwd._get_indexer_fwd_kernel(4, 4) + clean_first = tilelang_indexer_fwd._get_clean_logits_kernel() + clean_second = tilelang_indexer_fwd._get_clean_logits_kernel() + + assert first is second + assert third is not first + assert len(fwd_builds) == 2 + assert clean_first is clean_second + assert len(clean_builds) == 1 + + tilelang_indexer_bwd._tilelang_indexer_bwd_kernel_cache.clear() + bwd_builds = [] + + def fake_bwd_builder(*args, **kwargs): + bwd_builds.append((args, kwargs)) + return make_fake_kernel() + + monkeypatch.setattr(tilelang_indexer_bwd, "tl_indexer_bwd_impl", fake_bwd_builder) + bwd_first = tilelang_indexer_bwd._get_indexer_bwd_kernel(8, 4, 32) + bwd_second = tilelang_indexer_bwd._get_indexer_bwd_kernel(8, 4, 32) + bwd_third = tilelang_indexer_bwd._get_indexer_bwd_kernel(16, 4, 32) + + assert bwd_first is bwd_second + assert bwd_third is not bwd_first + assert len(bwd_builds) == 2 + assert bwd_builds[0][1]["num_threads"] == 32 + assert bwd_builds[1][1]["num_threads"] == 128 + + tilelang_sparse_mla_bwd._tilelang_sparse_mla_preprocess_kernel_cache.clear() + tilelang_sparse_mla_bwd._tilelang_sparse_mla_bwd_kernel_cache.clear() + tilelang_sparse_mla_bwd._tilelang_sparse_mla_postprocess_kernel_cache.clear() + monkeypatch.setattr( + tilelang_sparse_mla_bwd, "preprocess", lambda *args, **kwargs: make_fake_kernel() + ) + monkeypatch.setattr( + tilelang_sparse_mla_bwd, "bwd", lambda *args, **kwargs: make_fake_kernel() + ) + monkeypatch.setattr( + tilelang_sparse_mla_bwd, "postprocess", lambda *args, **kwargs: make_fake_kernel() + ) + + preprocess_first = tilelang_sparse_mla_bwd._get_preprocess_kernel(2, 4) + preprocess_second = tilelang_sparse_mla_bwd._get_preprocess_kernel(2, 4) + sparse_bwd_first = tilelang_sparse_mla_bwd._get_bwd_kernel(2, 512, 64, 32, 1, 0.5, 80) + sparse_bwd_second = tilelang_sparse_mla_bwd._get_bwd_kernel(2, 512, 64, 32, 1, 0.5, 80) + postprocess_first = tilelang_sparse_mla_bwd._get_postprocess_kernel(512, 64, 1) + postprocess_second = tilelang_sparse_mla_bwd._get_postprocess_kernel(512, 64, 1) + + assert preprocess_first is preprocess_second + assert sparse_bwd_first is sparse_bwd_second + assert postprocess_first is postprocess_second + finally: + tilelang_indexer_fwd._tilelang_indexer_fwd_kernel_cache.clear() + tilelang_indexer_fwd._tilelang_indexer_clean_logits_kernel_cache.clear() + tilelang_indexer_bwd._tilelang_indexer_bwd_kernel_cache.clear() + tilelang_sparse_mla_bwd._tilelang_sparse_mla_preprocess_kernel_cache.clear() + tilelang_sparse_mla_bwd._tilelang_sparse_mla_bwd_kernel_cache.clear() + tilelang_sparse_mla_bwd._tilelang_sparse_mla_postprocess_kernel_cache.clear() + + +def test_tilelang_utils_noop_jit_and_require_tilelang(monkeypatch): + def fn(): + return "ok" + + monkeypatch.setattr(tilelang_utils, "HAVE_TILELANG", False) + assert tilelang_utils._noop_jit(fn) is fn + assert tilelang_utils._noop_jit()(fn) is fn + assert tilelang_utils.tilelang_jit(fn) is fn + with pytest.raises(ImportError, match="TileLang is required"): + tilelang_utils.require_tilelang() + + +def test_compute_topk_target_chunk_sum_shared_and_per_head_paths(monkeypatch): + tilelang_dsa._DSA_SCRATCH_CACHE.clear() + monkeypatch.setattr(tilelang_dsa, "_DSA_SCRATCH_CACHE_TOTAL_BYTES", 0) + + query_h = torch.tensor( + [[[1.0, 0.0], [0.0, 1.0]], [[1.0, 1.0], [1.0, -1.0]]], requires_grad=True + ) + key_shared = torch.tensor([[1.0, 0.0], [0.0, 1.0], [1.0, 1.0]], requires_grad=True) + idx_seq = torch.tensor([[0, 1], [1, 2]], dtype=torch.int64) + valid_seq = torch.tensor([[True, True], [True, False]]) + + shared = tilelang_dsa._compute_topk_target_chunk_sum( + query_h=query_h, + key_shared=key_shared, + key_per_head=None, + s0=0, + s1=2, + idx_seq=idx_seq, + valid_seq=valid_seq, + softmax_scale=1.0, + head_chunk_size=1, + topk_chunk_size=1, + sk=3, + hn=2, + ) + + assert shared.shape == (2, 2) + assert not shared.requires_grad + assert torch.allclose(shared.sum(dim=-1), torch.tensor([2.0, 2.0]), atol=1e-6) + assert shared[1, 1] == 0 + + key_per_head = torch.stack((key_shared, key_shared + 1.0)).detach().requires_grad_(True) + per_head = tilelang_dsa._compute_topk_target_chunk_sum( + query_h=query_h, + key_shared=None, + key_per_head=key_per_head, + s0=0, + s1=2, + idx_seq=idx_seq, + valid_seq=valid_seq, + softmax_scale=1.0, + head_chunk_size=2, + topk_chunk_size=2, + sk=3, + hn=2, + ) + + assert per_head.shape == (2, 2) + assert not per_head.requires_grad + assert torch.allclose(per_head.sum(dim=-1), torch.tensor([2.0, 2.0]), atol=1e-6) + assert per_head[1, 1] == 0 + + +def test_tilelang_dsa_fused_hook_guard_paths(monkeypatch): + q = torch.empty(2, 1, 2, 4) + k = torch.empty(3, 1, 4) + weights = torch.empty(2, 1, 2) + starts = torch.zeros(2, dtype=torch.int32) + ends = torch.ones(2, dtype=torch.int32) + + monkeypatch.setattr(tilelang_dsa, "lighting_indexer_indices", None) + assert tilelang_dsa.fused_qk_topk_lighting(q, k, weights, 2, starts, ends, 1) is None + + monkeypatch.setattr(tilelang_dsa, "lighting_indexer_indices", lambda *args, **kwargs: None) + assert tilelang_dsa.fused_qk_topk_lighting(q.squeeze(1), k, weights, 2, starts, ends, 1) is None + assert tilelang_dsa.fused_qk_topk_lighting(q, k[:, :0], weights, 2, starts, ends, 1) is None + + monkeypatch.setattr(tilelang_dsa, "lighting_indexer", None) + query = torch.empty(2, 1, 2, 4) + key = torch.empty(3, 1, 1, 4) + assert ( + tilelang_dsa.fused_qk_topk_lighting_with_streaming_sparse_kl( + q, + k, + weights, + 2, + starts, + ends, + 1, + query, + key, + 1.0, + 0.1, + SimpleNamespace(tp=SimpleNamespace(size=lambda: 1)), + ) + is None + ) + + def fail_lighting_indexer(*_args, **_kwargs): + raise AssertionError("unsupported indexer head count should fall back before TileLang") + + monkeypatch.setattr(tilelang_dsa, "lighting_indexer", fail_lighting_indexer) + assert ( + tilelang_dsa.fused_qk_topk_lighting_with_streaming_sparse_kl( + torch.empty(2, 1, 7, 4), + k, + torch.empty(2, 1, 7), + 2, + starts, + ends, + 1, + torch.empty(2, 1, 2, 4), + key, + 1.0, + 0.1, + SimpleNamespace(tp=SimpleNamespace(size=lambda: 1)), + ) + is None + ) + + monkeypatch.setattr(tilelang_dsa, "SparseMLA", None) + topk_indices = torch.zeros(1, 2, 64, dtype=torch.int32) + assert tilelang_dsa.fused_sparse_mla_absorbed(query, key, topk_indices, 1.0, 512) is None + + class FakeSparseMLA: + @staticmethod + def apply(q_t, kv_t, idx_t, softmax_scale): + del q_t, kv_t, idx_t, softmax_scale + return torch.empty(2, 2, 128), torch.empty(2, 2) + + monkeypatch.setattr(tilelang_dsa, "SparseMLA", FakeSparseMLA) + assert ( + tilelang_dsa.fused_sparse_mla_absorbed(query.squeeze(1), key, topk_indices, 1.0, 512) + is None + ) + assert ( + tilelang_dsa.fused_sparse_mla_absorbed(query, key.squeeze(2), topk_indices, 1.0, 512) + is None + ) + assert tilelang_dsa.fused_sparse_mla_absorbed(query, key, topk_indices[:0], 1.0, 512) is None + assert tilelang_dsa.fused_sparse_mla_absorbed(query, key, topk_indices[:, :1], 1.0, 512) is None + assert ( + tilelang_dsa.fused_sparse_mla_absorbed(query, key[..., :3], topk_indices, 1.0, 512) is None + ) + assert ( + tilelang_dsa.fused_sparse_mla_absorbed(query, key, topk_indices[..., :63], 1.0, 512) is None + ) + + query_supported = torch.empty(2, 1, 3, 576) + key_supported = torch.empty(2, 1, 1, 576) + topk_supported = torch.zeros(1, 2, 64, dtype=torch.int32) + assert ( + tilelang_dsa.fused_sparse_mla_absorbed( + query_supported, key_supported, topk_supported, 1.0, 256 + ) + is None + ) + assert ( + tilelang_dsa.fused_sparse_mla_absorbed( + query_supported, key_supported, topk_supported[..., :63], 1.0, 512 + ) + is None + ) + + class FailSparseMLA: + @staticmethod + def apply(*_args, **_kwargs): + raise AssertionError( + "unsupported SparseMLA head count should fall back before TileLang" + ) + + monkeypatch.setattr(tilelang_dsa, "SparseMLA", FailSparseMLA) + assert ( + tilelang_dsa.fused_sparse_mla_absorbed( + torch.empty(2, 1, 96, 576), key_supported, topk_supported, 1.0, 512 + ) + is None + ) + + monkeypatch.setattr(tilelang_dsa, "SparseMLA", FakeSparseMLA) + assert ( + tilelang_dsa.fused_sparse_mla_absorbed( + query_supported, key_supported, topk_supported, 1.0, 512 + ) + is None + ) + + +def test_fused_qk_topk_lighting_sanitizes_mocked_tilelang_indices(monkeypatch): + q = torch.empty(3, 1, 2, 4, dtype=torch.bfloat16) + k = torch.empty(5, 1, 4, dtype=torch.bfloat16) + weights = torch.empty(3, 1, 2) + starts = torch.tensor([0, 2, 4], dtype=torch.int32) + ends = torch.tensor([2, 4, 5], dtype=torch.int32) + calls = [] + + def fake_lighting_indexer_indices( + index_q, index_k, index_w, starts_arg, ends_arg, index_topk, use_relu=True + ): + del index_k, index_w, index_topk + calls.append((tuple(index_q.shape), starts_arg.clone(), ends_arg.clone(), use_relu)) + return torch.stack((starts_arg, ends_arg), dim=-1).to(torch.int32) + + monkeypatch.setattr(tilelang_dsa, "lighting_indexer_indices", fake_lighting_indexer_indices) + + topk = tilelang_dsa.fused_qk_topk_lighting( + q, k, weights, index_topk=2, starts=starts, ends=ends, block_size=2, use_relu=False + ) + + assert torch.equal(topk, torch.tensor([[[0, -1], [2, -1], [4, -1]]], dtype=torch.int32)) + assert [call[0] for call in calls] == [(2, 2, 4), (1, 2, 4)] + assert all(call[3] is False for call in calls) + + +def test_fused_sparse_mla_absorbed_batches_mocked_tilelang_outputs(monkeypatch): + class FakeSparseMLA: + @staticmethod + def apply(q_t, kv_t, idx_t, softmax_scale): + assert q_t.shape == (2, 2, 16, 576) + assert kv_t.shape == (2, 2, 1, 576) + assert idx_t.shape == (2, 2, 1, 64) + assert softmax_scale == 0.25 + batch_sums = q_t.float().sum(dim=(1, 2, 3)).to(dtype=q_t.dtype) + out = batch_sums.view(q_t.size(0), 1, 1, 1).expand( + q_t.size(0), q_t.size(1), q_t.size(2), 512 + ) + lse = torch.zeros(q_t.size(0), q_t.size(1), q_t.size(2)) + return out, lse + + monkeypatch.setattr(tilelang_dsa, "SparseMLA", FakeSparseMLA) + query = torch.zeros(2, 2, 16, 576, dtype=torch.bfloat16) + query[:, 1].fill_(1.0) + key = torch.zeros(2, 2, 1, 576, dtype=torch.bfloat16) + topk_indices = torch.zeros(2, 2, 64, dtype=torch.int32) + + output = tilelang_dsa.fused_sparse_mla_absorbed( + query, key, topk_indices, softmax_scale=0.25, v_channels=512 + ) + + assert output.shape == (2, 2, 16, 512) + assert torch.equal(output[:, 0], torch.zeros_like(output[:, 0])) + assert torch.equal(output[:, 1], torch.full_like(output[:, 1], 18432.0)) + + +def test_fused_sparse_mla_absorbed_pads_small_head_count_without_gradient_leak(monkeypatch): + class FakeSparseMLA: + @staticmethod + def apply(q_t, kv_t, idx_t, softmax_scale): + assert q_t.shape == (1, 2, 16, 576) + assert kv_t.shape == (1, 2, 1, 576) + assert idx_t.shape == (1, 2, 1, 64) + assert softmax_scale == 0.25 + assert torch.count_nonzero(q_t[:, :, 8:]) == 0 + out = q_t[..., :512] + kv_t[..., :512] + return out, torch.zeros(q_t.shape[:-1], dtype=torch.float32) + + monkeypatch.setattr(tilelang_dsa, "SparseMLA", FakeSparseMLA) + query = torch.randn(2, 1, 8, 576, dtype=torch.bfloat16, requires_grad=True) + key = torch.randn(2, 1, 1, 576, dtype=torch.bfloat16, requires_grad=True) + topk_indices = torch.zeros(1, 2, 64, dtype=torch.int32) + + output = tilelang_dsa.fused_sparse_mla_absorbed( + query, key, topk_indices, softmax_scale=0.25, v_channels=512 + ) + + assert output is not None + assert output.shape == (2, 1, 8, 512) + output.float().sum().backward() + assert torch.equal(query.grad[..., :512], torch.ones_like(query.grad[..., :512])) + assert torch.count_nonzero(query.grad[..., 512:]) == 0 + assert torch.equal(key.grad[..., :512], torch.full_like(key.grad[..., :512], 8.0)) + assert torch.count_nonzero(key.grad[..., 512:]) == 0 + + +def test_streaming_sparse_kl_path_with_mocked_tilelang_indexer(monkeypatch): + q = torch.empty(2, 1, 2, 4, dtype=torch.bfloat16) + k = torch.empty(4, 1, 4, dtype=torch.bfloat16) + weights = torch.empty(2, 1, 2) + starts = torch.tensor([0, 0], dtype=torch.int32) + ends = torch.tensor([4, 4], dtype=torch.int32) + query = torch.empty(2, 1, 2, 4, dtype=torch.bfloat16) + key = torch.empty(4, 1, 1, 4, dtype=torch.bfloat16) + query_valid_rows = torch.tensor([[True, False]]) + + def fake_lighting_indexer( + index_q, + index_k, + index_w, + starts_arg, + ends_arg, + index_topk, + topk_indices=None, + use_relu=True, + ): + del index_k, index_w, starts_arg, ends_arg, topk_indices, use_relu + topk_scores = torch.zeros(index_q.size(0), index_topk) + topk = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)[: index_q.size(0)] + return topk_scores, topk + + def fake_compute_topk_target_chunk_sum(**kwargs): + idx_seq = kwargs["idx_seq"] + return torch.ones(idx_seq.shape, dtype=torch.float32, device=idx_seq.device) + + monkeypatch.setattr(tilelang_dsa, "lighting_indexer", fake_lighting_indexer) + monkeypatch.setattr(tilelang_dsa, "is_supported_indexer_bwd_head_count", lambda *_args: True) + monkeypatch.setattr( + tilelang_dsa, "_compute_topk_target_chunk_sum", fake_compute_topk_target_chunk_sum + ) + + topk, loss = tilelang_dsa.fused_qk_topk_lighting_with_streaming_sparse_kl( + q=q, + k=k, + weights=weights, + index_topk=2, + starts=starts, + ends=ends, + block_size=2, + query=query, + key=key, + softmax_scale=0.5, + loss_coeff=2.0, + pg_collection=SimpleNamespace(tp=SimpleNamespace(size=lambda: 1)), + query_valid_rows=query_valid_rows, + calculate_per_token_loss=False, + seq_chunk_size=1, + head_chunk_size=1, + topk_chunk_size=1, + use_relu=False, + ) + + assert torch.equal(topk, torch.tensor([[[0, 1], [2, 3]]], dtype=torch.int32)) + assert loss.item() == 0.0 + + +def test_streaming_sparse_kl_uses_fused_target_when_supported(monkeypatch): + q = torch.empty(2, 1, 2, 4, dtype=torch.bfloat16) + k = torch.empty(4, 1, 4, dtype=torch.bfloat16) + weights = torch.empty(2, 1, 2) + starts = torch.zeros(2, dtype=torch.int32) + ends = torch.full((2,), 4, dtype=torch.int32) + query = torch.empty(2, 1, 2, 4, dtype=torch.bfloat16) + key = torch.empty(4, 1, 1, 4, dtype=torch.bfloat16) + calls = [] + + def fake_lighting_indexer( + index_q, + index_k, + index_w, + starts_arg, + ends_arg, + index_topk, + topk_indices=None, + use_relu=True, + ): + del index_k, index_w, starts_arg, ends_arg, topk_indices, use_relu + topk_scores = torch.zeros(index_q.size(0), index_topk, requires_grad=True) + topk = torch.tensor([[0, 1], [2, 3]], dtype=torch.int32)[: index_q.size(0)] + return topk_scores, topk + + def fake_target(query_arg, key_arg, indices_arg, softmax_scale): + calls.append((query_arg, key_arg, indices_arg.clone(), softmax_scale)) + return torch.ones(indices_arg.shape, dtype=torch.float32) + + def fail_python_target(**_kwargs): + raise AssertionError("the PyTorch target path should not run") + + monkeypatch.setattr(tilelang_dsa, "lighting_indexer", fake_lighting_indexer) + monkeypatch.setattr(tilelang_dsa, "is_supported_indexer_bwd_head_count", lambda *_args: True) + monkeypatch.setattr(tilelang_dsa, "_can_use_fused_sparse_indexer_target", lambda *_args: True) + monkeypatch.setattr(tilelang_dsa, "sparse_indexer_target_interface", fake_target) + monkeypatch.setattr(tilelang_dsa, "_can_use_fused_sparse_indexer_kl", lambda *_args: False) + monkeypatch.setattr(tilelang_dsa, "_compute_topk_target_chunk_sum", fail_python_target) + + topk, loss = tilelang_dsa.fused_qk_topk_lighting_with_streaming_sparse_kl( + q=q, + k=k, + weights=weights, + index_topk=2, + starts=starts, + ends=ends, + block_size=2, + query=query, + key=key, + softmax_scale=0.5, + loss_coeff=2.0, + pg_collection=SimpleNamespace(tp=SimpleNamespace(size=lambda: 1)), + seq_chunk_size=2, + ) + + assert torch.equal(topk, torch.tensor([[[0, 1], [2, 3]]], dtype=torch.int32)) + assert loss.item() == 0.0 + assert len(calls) == 1 + assert calls[0][0].shape == (2, 2, 4) + assert calls[0][1].shape == (4, 4) + assert calls[0][3] == 0.5 + + +@pytest.mark.parametrize("heads", [48, 96]) +def test_fused_sparse_indexer_target_and_kl_match_reference(heads): + if not torch.cuda.is_available(): + pytest.skip("CUDA is required for TileLang indexer-loss tests") + if not tilelang_indexer_loss.HAVE_TILELANG: + pytest.skip("TileLang indexer-loss kernels are unavailable") + + torch.manual_seed(1234 + heads) + seq_len = 2 + key_len = 256 + topk = 256 + dim = 576 + softmax_scale = dim**-0.5 + query = torch.randn(seq_len, heads, dim, device="cuda", dtype=torch.bfloat16) + key = torch.randn(key_len, dim, device="cuda", dtype=torch.bfloat16) + topk_indices = torch.arange(topk, device="cuda", dtype=torch.int32).repeat(seq_len, 1) + topk_indices[1, -16:] = -1 + valid = topk_indices >= 0 + + target = tilelang_indexer_loss.sparse_indexer_target_interface( + query, key, topk_indices, softmax_scale + ) + safe_indices = topk_indices.clamp(min=0).to(torch.int64) + selected_key = key.index_select(0, safe_indices.reshape(-1)).view(seq_len, topk, dim) + reference_scores = ( + torch.einsum("shd,skd->shk", query.float(), selected_key.float()) * softmax_scale + ) + reference_scores = reference_scores.masked_fill(~valid.unsqueeze(1), float("-inf")) + reference_target = torch.softmax(reference_scores, dim=-1).masked_fill(~valid.unsqueeze(1), 0.0) + reference_target = reference_target.sum(dim=1) + torch.testing.assert_close(target, reference_target, rtol=2e-2, atol=2e-2) + + logits = torch.randn(seq_len, topk, device="cuda", dtype=torch.float32, requires_grad=True) + loss = tilelang_indexer_loss.SparseIndexerKLLoss.apply(target, logits, valid) + loss.backward() + + normalized_target = dsa_indexer_loss.normalize_indexer_target(reference_target) + reference_log_probs = dsa_masking.masked_log_softmax(logits.detach(), valid, dim=-1) + reference_loss = dsa_indexer_loss.indexer_kl_sum(normalized_target, reference_log_probs, valid) + reference_grad = ( + reference_log_probs.exp().masked_fill(~valid, 0.0) - normalized_target + ).masked_fill(~valid, 0.0) + torch.testing.assert_close(loss, reference_loss, rtol=2e-3, atol=2e-3) + torch.testing.assert_close(logits.grad, reference_grad, rtol=2e-3, atol=2e-3) + + +@pytest.mark.parametrize("dtype", [torch.float16, torch.float32]) +def test_tilelang_ops_decline_non_bfloat16_inputs(monkeypatch, dtype): + def fail_if_called(*_args, **_kwargs): + raise AssertionError("TileLang kernel should not run for non-BF16 inputs") + + class FailSparseMLA: + @staticmethod + def apply(*args): + return fail_if_called(*args) + + monkeypatch.setattr(tilelang_dsa, "lighting_indexer_indices", fail_if_called) + monkeypatch.setattr(tilelang_dsa, "lighting_indexer", fail_if_called) + monkeypatch.setattr(tilelang_dsa, "SparseMLA", FailSparseMLA) + + q_indexer = torch.zeros((1, 1, 1, 1), dtype=dtype) + k_indexer = torch.zeros((1, 1, 1), dtype=dtype) + weights = torch.zeros((1, 1, 1), dtype=dtype) + starts = torch.tensor([0], dtype=torch.int32) + ends = torch.tensor([1], dtype=torch.int32) + query = torch.zeros((1, 1, 1, 1), dtype=dtype) + key = torch.zeros((1, 1, 1, 1), dtype=dtype) + topk_indices = torch.zeros((1, 1, 1), dtype=torch.int32) + + assert ( + tilelang_dsa.fused_qk_topk_lighting(q_indexer, k_indexer, weights, 1, starts, ends, 128) + is None + ) + assert ( + tilelang_dsa.fused_qk_topk_lighting_with_streaming_sparse_kl( + q=q_indexer, + k=k_indexer, + weights=weights, + index_topk=1, + starts=starts, + ends=ends, + block_size=128, + query=query, + key=key, + softmax_scale=1.0, + loss_coeff=0.01, + pg_collection=object(), + ) + is None + ) + assert tilelang_dsa.fused_sparse_mla_absorbed(query, key, topk_indices, 1.0, 1) is None + + +@pytest.mark.parametrize("num_heads", [8, 64]) +def test_fused_sparse_mla_absorbed_accepts_thd_sentinels(num_heads): + if not torch.cuda.is_available(): + pytest.skip("CUDA is required for TileLang SparseMLA tests") + if tilelang_dsa.SparseMLA is None: + pytest.skip("TileLang SparseMLA kernel is unavailable") + + torch.manual_seed(1234) + torch.cuda.manual_seed(1234) + + # Match the sequence bucket so the kernel receives the original B=1 tensor views. + # This exercises the canonical batch-stride handling rather than hiding it with padding. + seqlen = 256 + dim = 576 + v_channels = 512 + topk = 64 + + query = torch.randn( + (seqlen, 1, num_heads, dim), dtype=torch.bfloat16, device="cuda", requires_grad=True + ) + key = torch.randn((seqlen, 1, 1, dim), dtype=torch.bfloat16, device="cuda", requires_grad=True) + topk_indices = torch.full((1, seqlen, topk), -1, dtype=torch.int32, device="cuda") + for row in range(1, seqlen): + valid = min(row, topk) + topk_indices[0, row, :valid] = torch.arange(valid, dtype=torch.int32, device="cuda") + + output = tilelang_dsa.fused_sparse_mla_absorbed( + query, key, topk_indices, softmax_scale=1.0 / math.sqrt(dim), v_channels=v_channels + ) + + assert output is not None + assert output.shape == (seqlen, 1, num_heads, v_channels) + assert torch.isfinite(output).all() + assert output[0].abs().max() == 0 + + output.float().square().mean().backward() + assert query.grad is not None + assert key.grad is not None + assert torch.isfinite(query.grad).all() + assert torch.isfinite(key.grad).all() + assert query.grad[0].abs().max() == 0 diff --git a/tests/unit_tests/transformer/moe/test_aux_loss.py b/tests/unit_tests/transformer/moe/test_aux_loss.py index bd88ab150c4..e4e113dc7b6 100644 --- a/tests/unit_tests/transformer/moe/test_aux_loss.py +++ b/tests/unit_tests/transformer/moe/test_aux_loss.py @@ -384,6 +384,79 @@ def test_seq_aux_loss(self, tp_size, ep_size, cp_size): torch.testing.assert_close(aux_loss, seq_aux_loss) torch.testing.assert_close(grad1, grad2) + @pytest.mark.internal + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + @pytest.mark.parametrize("with_padding", [False, True]) + @pytest.mark.parametrize( + "tp_size,ep_size,cp_size", [(8, 1, 1), (4, 2, 1), (1, 1, 8), (2, 1, 4), (2, 2, 2)] + ) + def test_seq_aux_loss_mbs_invariant_per_token_loss( + self, tp_size, ep_size, cp_size, with_padding + ): + """seq_aux_loss gradient must be invariant to MBS under --calculate-per-token-loss. + + The same global batch is processed as N micro-batches of size 1 (MBS=1) and as one + micro-batch of size N (MBS=N). Both cover the same tokens, so the finalize-time + 1/total_tokens normalization is an identical constant and the accumulated + router-weight aux gradients must match. Before the fix (valid_token_count dropped the + bsz factor), the MBS=N gradient is scaled by 1/N and the assertion fails. The padding + case additionally checks the correction uses valid (non-padded) token counts. + """ + Utils.initialize_model_parallel( + tensor_model_parallel_size=tp_size, + expert_tensor_parallel_size=ep_size, + context_parallel_size=cp_size, + ) + model_parallel_cuda_manual_seed(42) + clear_aux_losses_tracker() + + router = self.new_router( + moe_router_load_balancing_type="seq_aux_loss", + moe_aux_loss_coeff=1.0, + moe_router_dtype="fp64", + calculate_per_token_loss=True, + # fp32 weights so the MBS=1 gradient (accumulated over N backward passes) + # is not degraded by bf16 rounding relative to the single MBS=N backward. + params_dtype=torch.float32, + bf16=False, + tensor_model_parallel_size=tp_size, + expert_tensor_parallel_size=ep_size, + context_parallel_size=cp_size, + ).cuda() + + seq_len = 32 + num_seqs = 4 + with get_cuda_rng_tracker().fork(): + hidden_states = torch.randn( + (seq_len, num_seqs, router.config.hidden_size), + device=torch.device("cuda"), + dtype=torch.float32, + ) + padding_mask = None + if with_padding: + # True marks padding tokens (second half of each sequence). + padding_mask = torch.zeros((seq_len, num_seqs), dtype=torch.bool, device="cuda") + padding_mask[seq_len // 2 :, :] = True + + def run(indices): + pmask = None if padding_mask is None else padding_mask[:, indices] + scores, _ = router(hidden_states[:, indices, :].contiguous(), padding_mask=pmask) + scores.backward(torch.zeros_like(scores)) # isolate the aux-loss gradient + clear_aux_losses_tracker() + + # MBS=1: N micro-batches of size 1, accumulating the aux-loss gradient. + router.weight.grad = None + for b in range(num_seqs): + run(slice(b, b + 1)) + grad_mbs1 = router.weight.grad.clone() + + # MBS=N: a single micro-batch of size N. + router.weight.grad = None + run(slice(0, num_seqs)) + grad_mbsN = router.weight.grad.clone() + + torch.testing.assert_close(grad_mbs1, grad_mbsN) + @pytest.mark.internal @pytest.mark.skipif( not torch.cuda.is_available() or not HAVE_ROUTER_FUSION, diff --git a/tests/unit_tests/transformer/moe/test_grouped_mlp.py b/tests/unit_tests/transformer/moe/test_grouped_mlp.py index 09e26605a98..0b7c354eba7 100644 --- a/tests/unit_tests/transformer/moe/test_grouped_mlp.py +++ b/tests/unit_tests/transformer/moe/test_grouped_mlp.py @@ -62,6 +62,7 @@ def __init__( single_grouped_weight, single_grouped_bias=False, delay_wgrad_compute=False, + scale_bias=False, ): super().__init__() self.num_gemms = num_gemms @@ -74,6 +75,7 @@ def __init__( self.single_grouped_weight = single_grouped_weight self.single_grouped_bias = single_grouped_bias self.delay_wgrad_compute = delay_wgrad_compute + self.scale_bias = scale_bias def need_backward_dw(self): return False @@ -147,18 +149,21 @@ def register_forward_hook(self, hook): assert ops[0].weight1 is module.linear_fc1.weight1 assert ops[0].bias0 is module.linear_fc1.bias0 assert ops[0].bias1 is module.linear_fc1.bias1 + assert ops[0].scale_bias is False assert ops[1].glu_interleave_size == 16 assert ops[2].device == "meta" assert ops[2].weight is module.linear_fc2.weight + assert ops[2].scale_bias is False assert hasattr(ops, "forward_pre_hook") assert hasattr(ops, "forward_post_hook") -def test_fused_forward_caches_ops_and_forwards_expected_arguments(): +@pytest.mark.parametrize("fc2_bias", [False, True], ids=["no_fc2_bias", "fc2_bias"]) +def test_fused_forward_caches_ops_and_forwards_expected_arguments(fc2_bias): class FakeFusedOps: - def __call__(self, hidden_states, fc1_tokens, probs, fc2_tokens): - self.args = (hidden_states, fc1_tokens, probs, fc2_tokens) - return hidden_states + 1 + def __call__(self, *args): + self.args = args + return args[0] + 1 module = TEGroupedMLP.__new__(TEGroupedMLP) # `_fused_forward` calls `skip_routed_expert_padding(config)` (added by PR 4071), which @@ -174,6 +179,7 @@ def __call__(self, hidden_states, fc1_tokens, probs, fc2_tokens): delay_offload_until_cuda_graph=False, ) module._fused_ops = None + module.linear_fc2 = SimpleNamespace(use_bias=fc2_bias) fused_ops = FakeFusedOps() module._make_fused_ops = lambda: fused_ops hidden_states = torch.zeros(2, 4) @@ -188,6 +194,10 @@ def __call__(self, hidden_states, fc1_tokens, probs, fc2_tokens): assert fused_ops.args[1] is tokens_per_expert assert fused_ops.args[2] is probs assert fused_ops.args[3] is tokens_per_expert + if fc2_bias: + assert fused_ops.args[4] is probs + else: + assert len(fused_ops.args) == 4 def test_apply_bias_returns_input_unchanged_when_bias_is_none(): @@ -286,6 +296,43 @@ def fc2_hook(submodule, _inputs, _kwargs, output): torch.testing.assert_close(output, torch.full_like(output, 2)) +def test_make_fused_impl_pre_forward_hook_exposes_fsdp_main_grad_for_fused_wgrad(): + class FakeGroupedLinear(torch.nn.Module): + def __init__(self, *, fuse_wgrad_accumulation): + super().__init__() + self.fuse_wgrad_accumulation = fuse_wgrad_accumulation + self.weight = torch.nn.Parameter(torch.ones(2, 2)) + self.bias = torch.nn.Parameter(torch.zeros(2)) + + module = TEGroupedMLP.__new__(TEGroupedMLP) + torch.nn.Module.__init__(module) + module.linear_fc1 = FakeGroupedLinear(fuse_wgrad_accumulation=True) + module.linear_fc2 = FakeGroupedLinear(fuse_wgrad_accumulation=False) + + fc1_main_grad = torch.empty_like(module.linear_fc1.weight) + module.linear_fc1.weight.get_main_grad = lambda: fc1_main_grad + module.linear_fc1.weight.overwrite_main_grad = False + + existing_main_grad = torch.empty_like(module.linear_fc1.bias) + module.linear_fc1.bias.main_grad = existing_main_grad + module.linear_fc1.bias.get_main_grad = pytest.fail + module.linear_fc1.bias.overwrite_main_grad = False + + fc2_main_grad = torch.empty_like(module.linear_fc2.weight) + module.linear_fc2.weight.get_main_grad = lambda: fc2_main_grad + module.linear_fc2.weight.overwrite_main_grad = False + + hook = module._make_fused_impl_pre_forward_hook() + hook(object()) + + assert module.linear_fc1.weight.main_grad is fc1_main_grad + assert module.linear_fc1.weight.overwrite_main_grad is True + assert module.linear_fc1.bias.main_grad is existing_main_grad + assert module.linear_fc1.bias.overwrite_main_grad is True + assert getattr(module.linear_fc2.weight, "main_grad", None) is None + assert module.linear_fc2.weight.overwrite_main_grad is False + + def test_make_fused_ops_handles_single_grouped_weight_for_fc1(monkeypatch): class FakeGroupedLinear(torch.nn.Module): def __init__( @@ -301,6 +348,7 @@ def __init__( single_grouped_weight, single_grouped_bias=False, delay_wgrad_compute=False, + scale_bias=False, ): super().__init__() self.num_gemms = num_gemms @@ -313,6 +361,7 @@ def __init__( self.single_grouped_weight = single_grouped_weight self.single_grouped_bias = single_grouped_bias self.delay_wgrad_compute = delay_wgrad_compute + self.scale_bias = scale_bias def need_backward_dw(self): return False @@ -394,6 +443,7 @@ def register_forward_hook(self, hook): assert ops[2].weight1 is module.linear_fc2.weight1 assert ops[2].bias0 is module.linear_fc2.bias0 assert ops[2].bias1 is module.linear_fc2.bias1 + assert ops[2].scale_bias is True def _make_fake_te_namespace(): @@ -413,6 +463,7 @@ def __init__( single_grouped_weight, single_grouped_bias=False, delay_wgrad_compute=False, + scale_bias=False, ): super().__init__() self.num_gemms = num_gemms @@ -425,6 +476,7 @@ def __init__( self.single_grouped_weight = single_grouped_weight self.single_grouped_bias = single_grouped_bias self.delay_wgrad_compute = delay_wgrad_compute + self.scale_bias = scale_bias def need_backward_dw(self): return False @@ -670,6 +722,28 @@ def test_is_fused_impl_supported_requires_cutedsl_env(monkeypatch): assert module._is_fused_impl_supported() is False +def test_is_fused_impl_supported_requires_scaled_fc2_bias(monkeypatch): + fake_te, FakeGroupedLinear = _make_fake_te_namespace() + + class FakeGroupedLinearWithoutScaleBias(FakeGroupedLinear): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + + fake_te.pytorch.GroupedLinear = FakeGroupedLinearWithoutScaleBias + fake_te.pytorch.ops.GroupedLinear = FakeGroupedLinearWithoutScaleBias + monkeypatch.setattr(experts_module, "te", fake_te) + monkeypatch.setattr(experts_module, "HAVE_TE", True) + monkeypatch.setattr(experts_module, "is_te_min_version", lambda _: True) + _install_fake_te_ops_modules(monkeypatch, fake_te) + + module = _make_fused_impl_support_module( + FakeGroupedLinearWithoutScaleBias, activation_func=F.silu, gated_linear_unit=True + ) + module.linear_fc2.use_bias = True + + assert module._is_fused_impl_supported() is False + + @pytest.mark.parametrize( ( "use_fused_weighted_squared_relu", @@ -1089,6 +1163,91 @@ def test_gpu_make_fused_ops_constructs_with_real_te(self): experts.linear_fc2, f"weight{idx}" ) + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + @pytest.mark.internal + def test_gpu_fused_path_scales_fc2_bias(self): + """FC2 bias and its gradients must use the per-token router probability.""" + try: + from transformer_engine.pytorch.ops import GroupedLinear + except ImportError: + pytest.skip("TE op fuser API not available") + import inspect + + if "scale_bias" not in inspect.signature(GroupedLinear.__init__).parameters: + pytest.skip("Installed TE op fuser GroupedLinear lacks `scale_bias` support") + + Utils.destroy_model_parallel() + Utils.initialize_model_parallel(1, 1) + + tf_config = TransformerConfig( + num_layers=1, + hidden_size=self.hidden_size, + num_attention_heads=4, + num_moe_experts=self.num_experts, + use_cpu_initialization=False, + add_bias_linear=True, + gated_linear_unit=True, + activation_func=F.silu, + bias_activation_fusion=False, + bias_dropout_fusion=False, + bf16=True, + params_dtype=torch.bfloat16, + moe_router_load_balancing_type="sinkhorn", + moe_router_topk=1, + moe_grouped_gemm=True, + use_transformer_engine_op_fuser=True, + ) + _set_random_seed(seed_=123, data_parallel_random_init=False) + submodules = get_submodules( + get_gpt_layer_with_transformer_engine_submodules( + self.num_experts, moe_grouped_gemm=True + ).mlp + ) + layer = MoELayer(tf_config, submodules) + layer = Float16Module(layer.config, layer).module + layer.cuda() + experts = layer.experts + assert isinstance(experts, TEGroupedMLP) + + with torch.no_grad(): + for linear in (experts.linear_fc1, experts.linear_fc2): + for expert_idx in range(self.num_experts): + getattr(linear, f"weight{expert_idx}").zero_() + getattr(linear, f"bias{expert_idx}").zero_() + experts.linear_fc2.bias0.fill_(2.0) + experts.linear_fc2.bias1.fill_(4.0) + + hidden_states = torch.zeros( + 3, self.hidden_size, dtype=torch.bfloat16, device="cuda", requires_grad=True + ) + tokens_per_expert = torch.tensor([2, 1], dtype=torch.int32, device="cuda") + probs = torch.tensor( + [0.25, 0.5, 0.125], dtype=torch.bfloat16, device="cuda", requires_grad=True + ) + + output, _ = experts(hidden_states, tokens_per_expert, probs) + expected_output = torch.cat( + ( + probs[:2, None] * torch.full_like(output[:2], 2.0), + probs[2:, None] * torch.full_like(output[2:], 4.0), + ) + ) + torch.testing.assert_close(output, expected_output) + + output.sum().backward() + expected_prob_grad = ( + torch.tensor([2.0, 2.0, 4.0], dtype=torch.bfloat16, device="cuda") * self.hidden_size + ) + torch.testing.assert_close(probs.grad, expected_prob_grad) + torch.testing.assert_close( + experts.linear_fc2.bias0.grad, + torch.ones_like(experts.linear_fc2.bias0) * probs[:2].detach().sum(), + ) + torch.testing.assert_close( + experts.linear_fc2.bias1.grad, + torch.ones_like(experts.linear_fc2.bias1) * probs[2:].detach().sum(), + ) + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @pytest.mark.internal def test_gpu_fused_path_loss_decreases(self): diff --git a/tests/unit_tests/transformer/moe/test_paged_stashing.py b/tests/unit_tests/transformer/moe/test_paged_stashing.py index 967a8313d8b..7f876969b04 100644 --- a/tests/unit_tests/transformer/moe/test_paged_stashing.py +++ b/tests/unit_tests/transformer/moe/test_paged_stashing.py @@ -118,6 +118,7 @@ def __init__( moe_permute_fusion=kwargs.get("moe_permute_fusion", False), moe_flex_dispatcher_backend=kwargs.get("moe_flex_dispatcher_backend", None), moe_ncclep_static_shape=kwargs.get("moe_ncclep_static_shape", False), + moe_ncclep_zero_copy=kwargs.get("moe_ncclep_zero_copy", False), moe_grouped_gemm=kwargs.get("moe_grouped_gemm", False), moe_paged_stash=kwargs.get("moe_paged_stash", False), moe_expert_rank_capacity_factor=kwargs.get("moe_expert_rank_capacity_factor", None), @@ -190,6 +191,18 @@ def is_hybrid_ep_available(): return HAVE_HYBRIDEP +def is_nccl_ep_zero_copy_available(): + """Zero-copy needs the newer TE symm-mem APIs (symm_mem_alloc/is_symm_backed), absent in a plain + NCCL-EP build.""" + if not is_nccl_ep_available(): + return False + try: + from transformer_engine.pytorch.ep import is_symm_backed, symm_mem_alloc # noqa: F401 + except ImportError: + return False + return True + + def is_nccl_ep_available(): from megatron.core.transformer.moe.fused_a2a import HAVE_TE_EP @@ -454,11 +467,18 @@ def teardown_method(self, method): Utils.destroy_model_parallel() @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + # NCCL EP static-shape paged stashing aborts in dev CI with a pybind11 GIL dec_ref failure. + @pytest.mark.flaky_in_dev @pytest.mark.internal - def test_forward_backward_4_layers(self): - """Test paged stashing with 4 MoE layers on ncclep static shape: two passes match.""" + @pytest.mark.parametrize("zero_copy", [False, True]) + def test_forward_backward_4_layers(self, zero_copy): + """Test paged stashing with 4 MoE layers on ncclep static shape: two passes match. + + zero_copy=True additionally exercises the ncclEP symm-mem zero-copy IO under paged stash.""" if not is_nccl_ep_available(): pytest.skip("NCCL EP is not available") + if zero_copy and not is_nccl_ep_zero_copy_available(): + pytest.skip("NCCL EP zero-copy TE API is not available") config.ENABLE_EXPERIMENTAL = True @@ -485,6 +505,7 @@ def test_forward_backward_4_layers(self): moe_router_padding_for_quantization=True, gated_linear_unit=True, activation_func=F.silu, + moe_ncclep_zero_copy=zero_copy, ) seq_length = 1024 diff --git a/tests/unit_tests/transformer/moe/test_qb_routing.py b/tests/unit_tests/transformer/moe/test_qb_routing.py new file mode 100644 index 00000000000..5ad26487fb6 --- /dev/null +++ b/tests/unit_tests/transformer/moe/test_qb_routing.py @@ -0,0 +1,139 @@ +# Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. + +from typing import cast + +import pytest +import torch + +from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules +from megatron.core.transformer.moe.moe_layer import MoELayer, MoESubmodules +from megatron.core.transformer.moe.moe_utils import qb_dual_update +from megatron.core.transformer.moe.router import Router +from megatron.core.transformer.spec_utils import get_submodules +from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.training.initialize import _set_random_seed +from tests.unit_tests.test_utilities import Utils + + +class TestQBDualUpdate: + """Pure-tensor tests for the quantile-balancing dual update (CPU, no distributed).""" + + @pytest.mark.internal + @pytest.mark.parametrize("m,n,k", [(64, 8, 2), (40, 8, 1), (12, 4, 1)]) + def test_column_quantile_contract(self, m, n, k): + """qb_beta_local is the (col_target+1)-th largest score minus alpha per expert.""" + torch.manual_seed(123) + scores = torch.randn(m, n) + beta = torch.zeros(n) + + _, beta_local = qb_dual_update(scores, k, beta, update_beta=True) + + alpha = (scores - beta).topk(k + 1, dim=1).values[:, -1:] + adjusted = scores - alpha + col_target = m * k // n + expected = adjusted.sort(dim=0, descending=True).values[col_target] + torch.testing.assert_close(beta_local, expected) + + +class TestQuantileBalancingRouter: + def setup_method(self, method): + Utils.initialize_model_parallel(1, 1) + _set_random_seed(seed_=123, data_parallel_random_init=False) + self.num_moe_experts = 8 + self.transformer_config = TransformerConfig( + num_layers=2, + hidden_size=12, + num_attention_heads=4, + num_moe_experts=self.num_moe_experts, + use_cpu_initialization=True, + moe_router_load_balancing_type="quantile_balancing", + moe_router_score_function="softmax", + moe_router_topk=2, + moe_aux_loss_coeff=0, + bf16=True, + params_dtype=torch.bfloat16, + add_bias_linear=False, + ) + self.submodules = get_submodules( + get_gpt_layer_local_submodules( + num_experts=self.num_moe_experts, moe_grouped_gemm=False + ).mlp + ) + assert isinstance(self.submodules, MoESubmodules) + self.moe_layer = MoELayer(self.transformer_config, self.submodules) + self.router = cast(Router, self.moe_layer.router) + + def teardown_method(self, method): + Utils.destroy_model_parallel() + + @pytest.mark.internal + def test_non_qb_router_has_no_qb_buffers(self): + config = TransformerConfig( + num_layers=2, + hidden_size=12, + num_attention_heads=4, + num_moe_experts=self.num_moe_experts, + use_cpu_initialization=True, + moe_router_load_balancing_type="aux_loss", + moe_router_topk=2, + moe_aux_loss_coeff=0, + bf16=True, + params_dtype=torch.bfloat16, + add_bias_linear=False, + ) + router = MoELayer(config, self.submodules).router + assert router.qb_beta is None + assert router.qb_beta_accum is None + assert router.qb_beta_count is None + + @pytest.mark.internal + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + @pytest.mark.parametrize("moe_router_pre_softmax", [True, False]) + @pytest.mark.parametrize("score_function", ["softmax", "sigmoid"]) + def test_qb_router_forward(self, score_function, moe_router_pre_softmax): + self.router = self.router.cuda() + self.router.config.moe_router_score_function = score_function + self.router.score_function = score_function + self.router.config.moe_router_pre_softmax = moe_router_pre_softmax + + num_tokens = 32 * 2 + hidden_states = torch.randn((32, 2, self.router.config.hidden_size)).cuda().bfloat16() + with torch.no_grad(): + probs, routing_map = self.router(hidden_states) + + assert probs.shape == (num_tokens, self.num_moe_experts) + assert routing_map.shape == (num_tokens, self.num_moe_experts) + # Each token selects exactly topk distinct experts. + assert routing_map.sum().item() == num_tokens * self.router.topk + + @pytest.mark.internal + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + def test_qb_beta_accumulates_in_training(self): + self.router = self.router.cuda() + self.router.train() + hidden_states = torch.randn((32, 2, self.router.config.hidden_size)).cuda().bfloat16() + + assert self.router.qb_beta_count.item() == 0 + self.router(hidden_states) + assert self.router.qb_beta_count.item() == 1 + assert self.router.qb_beta_accum.abs().sum().item() > 0 + self.router(hidden_states) + assert self.router.qb_beta_count.item() == 2 + + # No accumulation outside the training path (eval / recompute). + accum_before = self.router.qb_beta_accum.clone() + with torch.no_grad(): + self.router(hidden_states) + assert self.router.qb_beta_count.item() == 2 + torch.testing.assert_close(self.router.qb_beta_accum, accum_before) + + @pytest.mark.internal + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + def test_qb_router_rejects_padding_mask(self): + self.router = self.router.cuda() + hidden_states = torch.randn((32, 2, self.router.config.hidden_size)).cuda().bfloat16() + padding_mask = torch.zeros((32, 2), dtype=torch.bool, device=hidden_states.device) + padding_mask[-2:] = True + + with pytest.raises(AssertionError, match="does not support padding masks"): + self.router(hidden_states, padding_mask=padding_mask) diff --git a/tests/unit_tests/transformer/moe/test_routers.py b/tests/unit_tests/transformer/moe/test_routers.py index b1fda34d483..cb4c7e60fb0 100644 --- a/tests/unit_tests/transformer/moe/test_routers.py +++ b/tests/unit_tests/transformer/moe/test_routers.py @@ -12,6 +12,7 @@ get_default_pg_collection, get_updated_expert_bias, router_gating_linear, + topk_routing_with_score_function, ) from megatron.core.transformer.moe.router import Router, TopKRouter from megatron.core.transformer.spec_utils import get_submodules @@ -627,6 +628,43 @@ def test_router_gating_linear_bias(router_dtype): assert torch.allclose(bias.grad, ref_bias.grad, **tols) +@pytest.mark.internal +@pytest.mark.parametrize("score_function", ["softmax", "sigmoid", "sqrtsoftplus"]) +@pytest.mark.parametrize("use_pre_softmax", [True, False]) +@pytest.mark.parametrize("topk", [1, 2]) +def test_topk_routing_precomputed_indices_equivalence(score_function, use_pre_softmax, topk): + """Passing precomputed_indices that match the function's own selection must reproduce + the standard output. Guards the shared post-top-k path reused by quantile balancing.""" + if score_function != "softmax" and use_pre_softmax: + pytest.skip("pre_softmax only applies to softmax scoring") + + torch.manual_seed(123) + num_tokens, num_experts = 64, 8 + logits = torch.randn(num_tokens, num_experts) + + kwargs = dict(use_pre_softmax=use_pre_softmax, score_function=score_function, fused=False) + probs_ref, map_ref = topk_routing_with_score_function(logits, topk, **kwargs) + _, top_indices = topk_routing_with_score_function(logits, topk, dense_output=True, **kwargs) + probs_pre, map_pre = topk_routing_with_score_function( + logits, topk, precomputed_indices=top_indices, **kwargs + ) + + # Natural top-k indices reproduce the standard output. + assert torch.equal(map_ref, map_pre) + torch.testing.assert_close(probs_ref, probs_pre) + + # Indices that differ from the natural top-k must route to exactly those experts. This + # catches a regression where the precomputed_indices branch is dropped and the function + # silently recomputes its own top-k instead of honoring the caller's indices. Bottom-k is + # disjoint from top-k since 2 * topk <= num_experts. + alt_indices = logits.topk(topk, dim=1, largest=False).indices + _, map_alt = topk_routing_with_score_function( + logits, topk, precomputed_indices=alt_indices, **kwargs + ) + expected_map = torch.zeros_like(logits, dtype=torch.bool).scatter(1, alt_indices, True) + assert torch.equal(map_alt, expected_map) + + # ============================================================ # Hash-based MoE routing tests # ============================================================ diff --git a/tests/unit_tests/transformer/moe/test_shared_experts.py b/tests/unit_tests/transformer/moe/test_shared_experts.py index 08da6c1ed0e..aedda040efe 100644 --- a/tests/unit_tests/transformer/moe/test_shared_experts.py +++ b/tests/unit_tests/transformer/moe/test_shared_experts.py @@ -10,12 +10,12 @@ from megatron.core.models.gpt import moe_module_specs from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules from megatron.core.parallel_state import get_tensor_model_parallel_world_size -from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed from megatron.core.transformer.moe import shared_experts as shared_experts_module from megatron.core.transformer.moe.moe_layer import MoELayer, MoESubmodules from megatron.core.transformer.moe.shared_experts import FusedSharedExpertMLP, SharedExpertMLP from megatron.core.transformer.spec_utils import get_submodules from megatron.core.transformer.transformer_config import TransformerConfig +from megatron.training.initialize import _set_random_seed from tests.unit_tests.test_utilities import Utils @@ -139,6 +139,8 @@ def _fake_shared_expert(**config_kwargs): shared_expert.tp_group = object() shared_expert._fused_grouped_swiglu_ops = None shared_expert._fused_grouped_swiglu_recipe = None + shared_expert._fused_grouped_swiglu_unit_scale = None + shared_expert._fused_grouped_swiglu_tokens_per_expert = {} return shared_expert @@ -227,6 +229,7 @@ def test_make_fused_grouped_swiglu_ops_builds_grouped_pipeline(monkeypatch): assert isinstance(activation_op, _FakeTEScaledSwiGLU) assert activation_op.glu_interleave_size == 32 + assert activation_op._grouped_mlp_unit_activation_scale is True assert isinstance(fc2_op, _FakeTEGroupedLinear) assert fc2_op.kwargs["num_groups"] == 1 @@ -279,12 +282,16 @@ def test_fused_grouped_swiglu_no_comm_flattens_and_caches_fused_ops(monkeypatch) (ops,) = shared_expert._fused_grouped_swiglu_ops hidden_states_2d, tokens_per_expert, scales, tokens_per_expert_again = ops.args + shared_expert._fused_grouped_swiglu_no_comm(torch.randn_like(hidden_states)) + _, cached_tokens_per_expert, cached_scales, _ = ops.args assert output.shape == hidden_states.shape assert shared_expert._fused_grouped_swiglu_recipe.__class__ is _FakeMXFP8Recipe assert hidden_states_2d.shape == (6, 4) assert tokens_per_expert.tolist() == [6] assert tokens_per_expert_again is tokens_per_expert - torch.testing.assert_close(scales, torch.ones(6)) + torch.testing.assert_close(scales, torch.ones(1)) + assert cached_tokens_per_expert is tokens_per_expert + assert cached_scales is scales def test_backward_dw_dispatches_fused_children_and_original_reduce_hooks(monkeypatch): @@ -353,13 +360,13 @@ def test_shared_expert_forward_backward(self, dispatcher_type: str, tp_size, ep_ tensor_model_parallel_size=tp_size, expert_model_parallel_size=ep_size ) # Create MoE layer with shared expert overlap enabled. - model_parallel_cuda_manual_seed(123) + _set_random_seed(seed_=123, data_parallel_random_init=False) moe_layer_overlap = self.get_moe_layer( moe_shared_expert_overlap=True, moe_token_dispatcher_type=dispatcher_type ).to(dtype=torch.bfloat16) # Create MoE layer with shared expert overlap disabled. - model_parallel_cuda_manual_seed(123) + _set_random_seed(seed_=123, data_parallel_random_init=False) moe_layer_no_overlap = self.get_moe_layer( moe_shared_expert_overlap=False, moe_token_dispatcher_type=dispatcher_type ).to(dtype=torch.bfloat16) diff --git a/tests/unit_tests/transformer/moe/test_token_dispatcher.py b/tests/unit_tests/transformer/moe/test_token_dispatcher.py index 42c43075ecf..b8f83d6d960 100644 --- a/tests/unit_tests/transformer/moe/test_token_dispatcher.py +++ b/tests/unit_tests/transformer/moe/test_token_dispatcher.py @@ -5,9 +5,14 @@ import pytest import torch +import torch.nn.functional as F from megatron.core import config, parallel_state -from megatron.core.models.gpt.gpt_layer_specs import get_gpt_layer_local_submodules +from megatron.core.fp8_utils import get_fp8_context +from megatron.core.models.gpt.gpt_layer_specs import ( + get_gpt_layer_local_submodules, + get_gpt_layer_with_transformer_engine_spec, +) from megatron.core.transformer.moe.fused_a2a import HYBRIDEP_TOKEN_ALIGNMENT, reset_hybrid_ep_buffer from megatron.core.transformer.moe.moe_layer import MoELayer, MoESubmodules from megatron.core.transformer.moe.moe_utils import get_capacity @@ -169,6 +174,13 @@ def __init__( moe_permute_fusion=kwargs.get("moe_permute_fusion", False), moe_flex_dispatcher_backend=kwargs.get("moe_flex_dispatcher_backend", None), moe_expert_rank_capacity_factor=kwargs.get("moe_expert_rank_capacity_factor", None), + moe_ncclep_static_shape=kwargs.get("moe_ncclep_static_shape", False), + moe_ncclep_zero_copy=kwargs.get("moe_ncclep_zero_copy", False), + use_transformer_engine_op_fuser=kwargs.get("use_transformer_engine_op_fuser", False), + gated_linear_unit=kwargs.get("gated_linear_unit", False), + activation_func=kwargs.get("activation_func", F.gelu), + fp8=kwargs.get("fp8", None), + fp8_recipe=kwargs.get("fp8_recipe", "delayed"), calculate_per_token_loss=kwargs.get("calculate_per_token_loss", False), ) @@ -176,14 +188,20 @@ def __init__( self.moe_layer = self.new_moe_layer() def new_moe_layer(self, **kargs): - submodules = get_submodules( - get_gpt_layer_local_submodules( + new_config = dataclasses.replace(self.config, **kargs) + if new_config.use_transformer_engine_op_fuser: + # op-fuser needs the TE grouped-MLP experts (they accept output_buffer/grad_input_buffer + # for the ncclEP zero-copy path); the local spec yields SequentialMLP, which does not. + mlp_spec = get_gpt_layer_with_transformer_engine_spec( + num_experts=new_config.num_moe_experts, moe_grouped_gemm=new_config.moe_grouped_gemm + ).submodules.mlp + else: + mlp_spec = get_gpt_layer_local_submodules( num_experts=self.config.num_moe_experts, moe_grouped_gemm=self.config.moe_grouped_gemm, ).mlp - ) + submodules = get_submodules(mlp_spec) assert isinstance(submodules, MoESubmodules) - new_config = dataclasses.replace(self.config, **kargs) moe_layer = MoELayer(new_config, submodules).cuda().to(dtype=self.test_dtype) moe_layer.set_layer_number(0) return moe_layer @@ -235,6 +253,52 @@ def dispatcher_dropless_test(self): hidden_states.grad, ans ), "Restored hidden states do not match original hidden states" + @pytest.mark.internal + def moe_layer_zero_copy_parity_test(self): + """Full MoE-layer fwd+bwd with ncclEP zero-copy OFF then ON (identical weights), asserting + parity. Runs the real op-fuser experts so fc2-out/fc1-dgrad are written straight into the + symm combine/dispatch buffers (verified via is_symm_backed) -- the pure permute/unpermute + harness cannot exercise this path.""" + from transformer_engine.pytorch.ep import is_symm_backed + + from megatron.core.transformer.moe.fused_a2a import nccl_ep_finalize + from megatron.core.transformer.moe.token_dispatcher import _NCCLEPManager + + torch.manual_seed(42) + x = torch.randn((32, 8, self.config.hidden_size), dtype=self.test_dtype).cuda() + + def run(layer): + inp = x.clone().detach().requires_grad_(True) + out, _ = layer(inp) # full fwd: dispatch -> op-fuser experts -> combine + out.sum().backward() # bwd: dispatch-bwd reads the symm grad buffer + return out.detach(), inp.grad.detach() + + def reset_ep(): + # zero_copy mode is fixed at ep_bootstrap (process-global); finalize + drop the shared + # symm classvars so the next layer re-bootstraps in the other mode. + nccl_ep_finalize() + _NCCLEPManager._zc_fwd_token_buf = None + _NCCLEPManager._zc_bwd_token_buf = None + _NCCLEPManager._zc_recv_topk_weights_buf = None + + ref_layer = self.new_moe_layer(moe_ncclep_zero_copy=False) + out_ref, grad_ref = run(ref_layer) + + reset_ep() + zc_layer = self.new_moe_layer(moe_ncclep_zero_copy=True) + zc_layer.load_state_dict(ref_layer.state_dict()) # identical weights + out_zc, grad_zc = run(zc_layer) + + # the combine forward buffer must be an allocated, registered symm window (zero-copy engaged) + fwd_buf = _NCCLEPManager._zc_fwd_token_buf + assert fwd_buf is not None, "zero-copy forward symm buffer was not allocated" + assert is_symm_backed(fwd_buf), "zero-copy forward buffer is not symm-mem-backed" + reset_ep() + + assert not torch.isnan(out_zc).any() and not torch.isnan(grad_zc).any() + torch.testing.assert_close(out_zc, out_ref, rtol=1e-2, atol=1e-2) + torch.testing.assert_close(grad_zc, grad_ref, rtol=1e-2, atol=1e-2) + @pytest.mark.internal def dispatcher_capacity_test(self): moe_layer = self.moe_layer @@ -518,6 +582,70 @@ def skip_if_flex_backend_unavailable(moe_flex_dispatcher_backend): pytest.skip("NCCL EP is not available") +def is_nccl_ep_zero_copy_available(): + """Zero-copy needs the newer TE symm-mem APIs (symm_mem_alloc/is_symm_backed), which a plain + NCCL-EP build lacks -- gate zero-copy tests on these separately from is_nccl_ep_available().""" + if not is_nccl_ep_available(): + return False + try: + from transformer_engine.pytorch.ep import is_symm_backed, symm_mem_alloc # noqa: F401 + except ImportError: + return False + return True + + +def is_op_fuser_available(): + """The static-shape/zero-copy path runs the TE op-fuser grouped GEMM (needs TE>=2.14 ops).""" + try: + from transformer_engine.pytorch.ops import GroupedLinear, ScaledSwiGLU # noqa: F401 + except ImportError: + return False + return is_te_min_version("2.14.0") + + +def test_hybridep_pad_uneven_dispatch_inputs_metadata(monkeypatch): + manager = _HybridEPManager.__new__(_HybridEPManager) + manager.group = object() + manager.num_local_experts = 2 + manager.num_experts = 4 + manager.config = TransformerConfig( + num_layers=1, + hidden_size=16, + num_attention_heads=4, + num_moe_experts=4, + moe_router_topk=2, + moe_hybridep_pad_variable_tokens=True, + ) + manager.moe_expert_rank_capacity_factor = None + manager.drop_and_pad = False + + local_num_tokens = 17 + max_num_tokens_across_ep = 70 + padded_num_tokens = ( + max_num_tokens_across_ep + -max_num_tokens_across_ep % HYBRIDEP_TOKEN_ALIGNMENT + ) + routing_map = torch.ones((local_num_tokens, manager.num_experts), dtype=torch.bool) + probs = torch.ones((local_num_tokens, manager.num_experts), dtype=torch.float32) + + def fake_all_reduce(tensor, op=None, group=None): + assert op == torch.distributed.ReduceOp.MAX + assert group is manager.group + tensor.fill_(max_num_tokens_across_ep) + + monkeypatch.setattr(torch.distributed, "all_reduce", fake_all_reduce) + + manager.setup_metadata(routing_map, probs) + + assert manager._original_num_tokens == local_num_tokens + assert manager._padded_num_tokens == padded_num_tokens + assert manager.routing_map.shape == (padded_num_tokens, manager.num_experts) + assert manager.token_probs.shape == (padded_num_tokens, manager.num_experts) + torch.testing.assert_close(manager.routing_map[:local_num_tokens], routing_map) + torch.testing.assert_close(manager.token_probs[:local_num_tokens], probs) + assert not manager.routing_map[local_num_tokens:].any() + assert not manager.token_probs[local_num_tokens:].any() + + @pytest.mark.skipif( not is_deep_ep_available() and not is_deep_ep_v2_available() and not is_hybrid_ep_available(), reason="Deep EP, Deep EP v2 and Hybrid EP are not available", @@ -535,7 +663,14 @@ def teardown_method(self, method): @pytest.mark.parametrize("tp_size,ep_size", [(1, 8), (8, 1), (4, 2)]) @pytest.mark.parametrize("permute_fusion", permute_fusion_params) @pytest.mark.parametrize( - "moe_flex_dispatcher_backend", ["deepep", "deepepv2", "hybridep", "ncclep"] + "moe_flex_dispatcher_backend", + [ + "deepep", + "deepepv2", + "hybridep", + # NCCL EP aborts in dev CI with a pybind11 GIL dec_ref failure. + pytest.param("ncclep", marks=pytest.mark.flaky_in_dev), + ], ) @pytest.mark.parametrize("moe_permute_fusion_into_hybridep", [True, False]) def test_forward_backward( @@ -579,6 +714,41 @@ def test_forward_backward( # reset experimental flag to False config.ENABLE_EXPERIMENTAL = False + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") + @pytest.mark.skipif( + not is_nccl_ep_zero_copy_available(), reason="NCCL EP zero-copy TE API is not available" + ) + @pytest.mark.skipif( + not is_op_fuser_available(), reason="op-fuser (static-shape/zero-copy) needs TE>=2.14" + ) + @pytest.mark.internal + @pytest.mark.timeout(120) + @pytest.mark.parametrize("tp_size,ep_size", [(1, 8)]) + def test_forward_backward_zero_copy(self, tp_size, ep_size): + # zero-copy requires static_shape, which requires BOTH op-fuser and grouped_gemm; bf16 so no + # fp8/Blackwell dependency. The op-fuser needs tp=1 and a SwiGLU activation. Parity: the + # zero-copy IO path must match the staged (no-zc) path. + container = MoEModelTestContainer( + tp_size=tp_size, + ep_size=ep_size, + pp_size=1, + num_moe_experts=8, + moe_router_topk=2, + moe_router_load_balancing_type="aux_loss", + moe_token_dispatcher_type="flex", + moe_flex_dispatcher_backend="ncclep", + moe_grouped_gemm=True, + use_transformer_engine_op_fuser=True, + moe_ncclep_static_shape=True, + gated_linear_unit=True, + activation_func=F.silu, + # ncclep sizes a per-rank recv buffer from this and overflow HARD-TRAPS; size generously. + moe_expert_rank_capacity_factor=8.0, + hidden_size=1024, + test_dtype=torch.bfloat16, + ) + container.moe_layer_zero_copy_parity_test() + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available") @pytest.mark.internal @pytest.mark.timeout(120) diff --git a/tests/unit_tests/transformer/test_cuda_graphs.py b/tests/unit_tests/transformer/test_cuda_graphs.py index 8ee5c34e2dd..8d53f3cbc54 100644 --- a/tests/unit_tests/transformer/test_cuda_graphs.py +++ b/tests/unit_tests/transformer/test_cuda_graphs.py @@ -23,6 +23,7 @@ destroy_num_microbatches_calculator, init_num_microbatches_calculator, ) +from megatron.core.packed_seq_params import PackedSeqParams from megatron.core.pipeline_parallel.schedules import set_current_microbatch from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import ( @@ -35,8 +36,14 @@ TECudaGraphHelper, _CudagraphGlobalRecord, _layer_is_graphable, + create_cudagraphs, +) +from megatron.core.transformer.enums import ( + AttnBackend, + CudaGraphModule, + CudaGraphScope, + InferenceCudaGraphScope, ) -from megatron.core.transformer.enums import CudaGraphModule, CudaGraphScope, InferenceCudaGraphScope from megatron.core.transformer.mlp import MLPSubmodules from megatron.core.transformer.module import GraphableMegatronModule, MegatronModule from megatron.core.transformer.moe.fused_a2a import reset_hybrid_ep_buffer @@ -446,6 +453,191 @@ def test_gpu_cudagraph(self): ) +@pytest.mark.skipif( + not (HAVE_TE and is_te_min_version("1.5.0")), + reason="use_te_rng_tracker requires TransformerEngine version >= 1.5", +) +class TestPackedSeqCudagraphs: + """Training CUDA graphs over thd input with padding between sequences. + + The padded cu_seqlens describe a slot layout that differs from the actual lengths, + and pad_between_seqs is set explicitly so TE does spend a GPU sync inferring it. + cp_size == 2 additionally captures TE's ring-P2P context-parallel attention inside the graphs. + """ + + SEQ_LENGTHS = [7, 5] + SLOT_STARTS = [0, 8, 16] # slot layout aligned to 2 * cp_size for every cp_size tested + BIN_SIZE = 32 + NVTE_ENV_VARS = ( + "NVTE_FLASH_ATTN", + "NVTE_FUSED_ATTN", + "NVTE_UNFUSED_ATTN", + "NVTE_ALLOW_NONDETERMINISTIC_ALGO", + ) + + def setup_method(self, method): + self.original_nvte_env = {name: os.environ.get(name) for name in self.NVTE_ENV_VARS} + os.environ["NVTE_ALLOW_NONDETERMINISTIC_ALGO"] = "0" + + def teardown_method(self, method): + try: + Utils.destroy_model_parallel() + _CudagraphGlobalRecord.cudagraph_created = False + _CudagraphGlobalRecord.cudagraph_record = [] + CudaGraphManager.global_mempool = None + finally: + for name, value in self.original_nvte_env.items(): + if value is None: + os.environ.pop(name, None) + else: + os.environ[name] = value + + def _build_packed_seq_params(self, device): + # Actual boundaries: each sequence's real tokens inside its slot; the trailing bin + # padding [SLOT_STARTS[-1], BIN_SIZE) forms a ghost slot of pad tokens. + boundaries = [0] + for length in self.SEQ_LENGTHS: + boundaries.append(boundaries[-1] + length) + boundaries.append(boundaries[-1] + self.BIN_SIZE - self.SLOT_STARTS[-1]) + cu_seqlens = torch.tensor(boundaries, dtype=torch.int32, device=device) + cu_seqlens_padded = torch.tensor( + self.SLOT_STARTS + [self.BIN_SIZE], dtype=torch.int32, device=device + ) + return PackedSeqParams( + qkv_format='thd', + cu_seqlens_q=cu_seqlens, + cu_seqlens_kv=cu_seqlens, + cu_seqlens_q_padded=cu_seqlens_padded, + cu_seqlens_kv_padded=cu_seqlens_padded, + max_seqlen_q=self.BIN_SIZE, + max_seqlen_kv=self.BIN_SIZE, + pad_between_seqs=True, + ) + + @pytest.mark.parametrize("cp_size", [1, 2]) + def test_thd_capture_with_pad_between_seqs(self, cp_size): + initialize_rng_tracker(use_te_rng_tracker=True, force_reset=True) + Utils.initialize_model_parallel(context_parallel_size=cp_size) + model_parallel_cuda_manual_seed(123) + os.environ["NVTE_FLASH_ATTN"] = "0" + os.environ["NVTE_FUSED_ATTN"] = "1" + os.environ["NVTE_UNFUSED_ATTN"] = "0" + + config = TransformerConfig( + num_layers=2, + hidden_size=64, + num_attention_heads=4, + context_parallel_size=cp_size, + bf16=True, + params_dtype=torch.bfloat16, + attention_dropout=0.0, + hidden_dropout=0.0, + attention_backend=AttnBackend.fused, + deterministic_mode=True, + cuda_graph_impl="local", + cuda_graph_warmup_steps=1, + use_cpu_initialization=True, + ) + block = TransformerBlock(config, get_gpt_layer_with_transformer_engine_spec()).cuda() + block.train() + # CUDA-graphed backward assumes DDP-style grad accumulation buffers. + for param in block.parameters(): + param.main_grad = torch.zeros_like(param) + + packed_seq_params = self._build_packed_seq_params(torch.device('cuda')) + # Each CP rank holds its 1/cp_size share of the bin's tokens. + hidden_states = torch.randn( + (self.BIN_SIZE // cp_size, 1, config.hidden_size), + dtype=torch.bfloat16, + device='cuda', + requires_grad=True, + ) + + eager_out = block( + hidden_states=hidden_states, attention_mask=None, packed_seq_params=packed_seq_params + ) + hidden_states_metadata = hidden_states.cg_buffer_metadata + assert hidden_states_metadata.is_cudagraph_input + assert hidden_states_metadata.is_saved_for_backward + + # The second layer's TE input layernorm saves the first layer's output for backward. + # This naturally exercises a CUDA graph output whose pool buffer must stay alive until + # backward capture. + first_runner = block.layers[0].cudagraph_manager.cudagraph_runners[0] + first_runner_record = next( + record + for record in _CudagraphGlobalRecord.cudagraph_record + if record[0] is first_runner and record[1] == "fwd" + ) + recorded_outputs = first_runner_record[4] + output_metadata = first_runner.get_arg_metas(recorded_outputs)[0].cg_buffer_metadata + output_metadata_state = ( + f"input={output_metadata.is_cudagraph_input}, " + f"output={output_metadata.is_cudagraph_output}, " + f"saved={output_metadata.is_saved_for_backward}" + ) + assert output_metadata.is_cudagraph_input, output_metadata_state + assert output_metadata.is_cudagraph_output, output_metadata_state + assert output_metadata.is_saved_for_backward, output_metadata_state + + # The q/kv aliases for each offsets tensor must share one metadata object while recording + # every graph-input use for replay-buffer sharing. + actual_cu_seqlens_metadata = packed_seq_params.cu_seqlens_q.cg_buffer_metadata + padded_cu_seqlens_metadata = packed_seq_params.cu_seqlens_q_padded.cg_buffer_metadata + assert packed_seq_params.cu_seqlens_kv.cg_buffer_metadata is actual_cu_seqlens_metadata + assert ( + packed_seq_params.cu_seqlens_kv_padded.cg_buffer_metadata is padded_cu_seqlens_metadata + ) + assert actual_cu_seqlens_metadata.is_cudagraph_input + assert padded_cu_seqlens_metadata.is_cudagraph_input + eager_out.sum().backward() + + # This is the primary function under test. + create_cudagraphs() + + runners = [] + for layer in block.layers: + layer_runners = layer.cudagraph_manager.cudagraph_runners + assert len(layer_runners) == 1 + assert layer_runners[0].fwd_graph is not None + runners.extend(layer_runners) + + # There are four cu_seqlens arguments per layer: q/kv pairs for the real and padded + # offsets. Each pair and every later layer should alias one of two shared buffers. Within + # each buffer group, only its first graph-input occurrence performs the replay copy. + cu_seqlens_buffers = [ + tensor + for runner in runners + for tensor in runner.fwd_graph_input_surface[: runner.num_dgrads] + if tensor.dtype == torch.int32 and tensor.shape == packed_seq_params.cu_seqlens_q.shape + ] + assert len(cu_seqlens_buffers) == 4 * len(runners) + buffers_by_ptr = {} + for tensor in cu_seqlens_buffers: + buffers_by_ptr.setdefault(tensor.data_ptr(), []).append(tensor) + assert len(buffers_by_ptr) == 2 + for shared_buffers in buffers_by_ptr.values(): + assert sum(not tensor.can_skip_replay_copy for tensor in shared_buffers) == 1 + + graphed_out = block( + hidden_states=hidden_states, attention_mask=None, packed_seq_params=packed_seq_params + ) + assert torch.equal(graphed_out, eager_out), ( + "CUDA graph replay output is not bitwise equal to eager output: " + f"max_abs_diff={(graphed_out.float() - eager_out.float()).abs().max().item()}" + ) + graphed_out.sum().backward() + + # Destroy captured graphs deterministically before parallel-state teardown. + for layer in block.layers: + for runner in layer.cudagraph_manager.cudagraph_runners: + if hasattr(runner, "fwd_graph"): + del runner.fwd_graph + if hasattr(runner, "bwd_graph"): + del runner.bwd_graph + torch.cuda.synchronize() + + @pytest.mark.skipif( not (HAVE_TE and is_te_min_version("1.5.0")), reason="use_te_rng_tracker requires TransformerEngine version >= 1.5", diff --git a/tests/unit_tests/transformer/test_full_cuda_graph.py b/tests/unit_tests/transformer/test_full_cuda_graph.py index 312ae467304..037b9dde287 100644 --- a/tests/unit_tests/transformer/test_full_cuda_graph.py +++ b/tests/unit_tests/transformer/test_full_cuda_graph.py @@ -1,23 +1,95 @@ # Copyright (c) 2025, NVIDIA CORPORATION. All rights reserved. +import warnings +from unittest.mock import Mock, patch + import pytest import torch from pytest_mock import mocker import megatron.core.pipeline_parallel.schedules as schedule from megatron.core import ModelParallelConfig -from megatron.core.full_cuda_graph import FullCudaGraphWrapper +from megatron.core.full_cuda_graph import FullCudaGraphWrapper, get_shared_capture_stream from megatron.core.tensor_parallel.random import ( HAVE_TE, initialize_rng_tracker, model_parallel_cuda_manual_seed, ) from megatron.core.utils import is_te_min_version +from megatron.training.models.dist_utils import _ddp_wrap from tests.unit_tests.test_utilities import Utils rank = Utils.rank +def test_ddp_grad_accumulators_share_full_cuda_graph_stream(): + """Retained DDP AccumulateGrad nodes must use the full-iteration capture stream.""" + + class RetainingDataParallel(torch.nn.Module): + """Minimal DDP wrapper that retains parameter AccumulateGrad nodes.""" + + def __init__(self, *, module, **_): + super().__init__() + self.module = module + self.grad_accumulators = [] + for param in module.parameters(): + expanded_param = param.expand_as(param) + grad_accumulator = expanded_param.grad_fn.next_functions[0][0] + grad_accumulator.register_hook(lambda *_: None) + self.grad_accumulators.append(grad_accumulator) + + def forward(self, inputs): + """Run the wrapped module.""" + return self.module(inputs) + + assert torch.autograd.graph.set_warn_on_accumulate_grad_stream_mismatch is not None + model = torch.nn.Linear(4, 4, device="cuda") + model.config = Mock(cuda_graph_impl="full_iteration") + ddp_config = Mock( + num_buckets=None, + bucket_size=1024, + overlap_grad_reduce=True, + use_distributed_optimizer=False, + ) + process_groups = Mock() + with patch( + "megatron.training.models.dist_utils.DistributedDataParallel", RetainingDataParallel + ): + wrapped_model = _ddp_wrap( + [model], + data_parallel_random_init=False, + ddp_config=ddp_config, + overlap_param_gather_with_optimizer_step=False, + pg_collection=process_groups, + )[0] + + capture_stream = get_shared_capture_stream() + current_stream = torch.cuda.current_stream() + capture_stream.wait_stream(current_stream) + static_input = torch.ones(2, 4, device="cuda") + + with warnings.catch_warnings(record=True) as caught_warnings: + warnings.simplefilter("always") + with torch.cuda.stream(capture_stream): + wrapped_model(static_input).sum().backward() + wrapped_model.zero_grad(set_to_none=False) + + cuda_graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(cuda_graph, stream=capture_stream): + wrapped_model(static_input).sum().backward() + + cuda_graph.replay() + torch.cuda.synchronize() + + stream_mismatch_warnings = [ + warning + for warning in caught_warnings + if "AccumulateGrad node's stream does not match" in str(warning.message) + ] + assert not stream_mismatch_warnings + assert all(param.grad is not None for param in wrapped_model.parameters()) + + @pytest.mark.skipif( not (HAVE_TE and is_te_min_version("1.5.0")), reason="use_te_rng_tracker requires TransformerEngine version >= 1.5", diff --git a/tests/unit_tests/transformer/test_submodule_callables.py b/tests/unit_tests/transformer/test_submodule_callables.py index ecbefe50d63..7b66f7d53d2 100644 --- a/tests/unit_tests/transformer/test_submodule_callables.py +++ b/tests/unit_tests/transformer/test_submodule_callables.py @@ -2,7 +2,8 @@ import pytest import torch -from megatron.core.models.gpt.fine_grained_callables import build_layer_callables +from megatron.core.models.common import fine_grained_callables as common_callables +from megatron.core.models.common.fine_grained_callables import build_layer_callables from megatron.core.models.gpt.gpt_layer_specs import ( get_gpt_layer_with_transformer_engine_submodules, ) @@ -15,6 +16,7 @@ compare_captures, deterministic_mode, get_test_config, + get_valid_flex_dispatcher_backend, get_valid_token_dispatcher_types, reset_model, ) @@ -101,6 +103,71 @@ def run_model_submodules_with_capture(model, input_tensors, microbatches): return capture +def test_mtp_pre_dispatch_applies_hybrid_empty_decoder_final_norm(monkeypatch): + """Covers the HybridModel empty-decoder MTP pre-dispatch final_norm path.""" + + from megatron.core.models.hybrid.hybrid_model import HybridModel + + def inner_pre_dispatch(_node, hidden_states): + return hidden_states + + def unused_forward(*_args, **_kwargs): + raise AssertionError("only MTP pre-dispatch should run in this test") + + def fake_build_layer_callables(_layer): + return ( + [inner_pre_dispatch, unused_forward, unused_forward, unused_forward, None], + {"pre_dispatch_computation": object()}, + ) + + class FakeMTPConfig: + sequence_parallel = False + + class FakeMTPLayer: + config = FakeMTPConfig() + eh_proj = object() + mtp_model_layer = object() + + def _get_embeddings( + self, input_ids, position_ids, embedding, hidden_states, packed_seq_params, padding_mask + ): + return input_ids, position_ids, padding_mask, None, hidden_states + + def _concat_embeddings(self, hidden_states, decoder_input): + return hidden_states + + def _postprocess(self, hidden_states): + return hidden_states + + monkeypatch.setattr(common_callables, "build_layer_callables", fake_build_layer_callables) + monkeypatch.setattr(common_callables, "get_layer_moe_metadata", lambda _layer: (True, 1)) + monkeypatch.setattr(common_callables, "get_mtp_layer_offset", lambda _config, _vp_stage: 0) + + model = HybridModel.__new__(HybridModel) + torch.nn.Module.__init__(model) + model.decoder = DummyState() + model.decoder.layers = [] + model.decoder.final_norm = lambda hidden_states: hidden_states + 4.0 + model.embedding = object() + model.vp_stage = None + + node = DummyNode() + node.chunk_state = DummyState() + node.chunk_state.model = model + node.chunk_state.context = None + node.chunk_state.packed_seq_params = None + node.is_first_layer = True + + hidden_states = torch.arange(6, dtype=torch.float32).reshape(3, 1, 2).requires_grad_() + expected = hidden_states + 4.0 + forward_funcs, _ = common_callables.build_mtp_layer_callables(FakeMTPLayer()) + + output = forward_funcs[0](node, hidden_states) + + torch.testing.assert_close(output, expected) + torch.testing.assert_close(node.chunk_state.mtp_hidden_states[0], expected) + + class TestTransformerLayerSubmoduleCallables: """ Test class for transformer layer submodule callables. @@ -133,19 +200,21 @@ def test_1f1b_overlap(self, dispatcher_type, grouped_gemm, permute_fusion): expert_model_parallel_size=2, virtual_pipeline_model_parallel_size=2, ) + qk_layernorm = True extra_kwargs = { "moe_token_dispatcher_type": dispatcher_type, "moe_permute_fusion": permute_fusion, + "qk_layernorm": qk_layernorm, } if dispatcher_type == "flex": - extra_kwargs["moe_flex_dispatcher_backend"] = "deepep" + extra_kwargs["moe_flex_dispatcher_backend"] = get_valid_flex_dispatcher_backend() config = get_test_config(extra_kwargs=extra_kwargs, moe_grouped_gemm=grouped_gemm) microbatches = 4 with deterministic_mode(): transformer_layer_submodules = get_gpt_layer_with_transformer_engine_submodules( num_experts=8, moe_grouped_gemm=grouped_gemm, - qk_layernorm=True, + qk_layernorm=qk_layernorm, multi_latent_attention=True, ) model = TransformerLayer(config, transformer_layer_submodules) diff --git a/tests/unit_tests/transformer/test_te_layers_batch_invariant.py b/tests/unit_tests/transformer/test_te_layers_batch_invariant.py index e2d52727925..685e9332025 100644 --- a/tests/unit_tests/transformer/test_te_layers_batch_invariant.py +++ b/tests/unit_tests/transformer/test_te_layers_batch_invariant.py @@ -29,6 +29,16 @@ except ImportError: HAVE_FA3 = False +try: + from flash_attn.cute import flash_attn_varlen_func as _fa4_varlen_func # noqa: F401 + + HAVE_FA4 = True +except ImportError: + HAVE_FA4 = False + +# Batch-invariant mode requires an explicit FlashAttention version. +_BIK_FA_VERSION = 4 if HAVE_FA4 else 3 + # ============================================================================ # Batch-Invariant test helpers @@ -108,6 +118,7 @@ def test_te_column_parallel_linear_batch_invariant_randomized(): hidden_dropout=0.0, attention_dropout=0.0, batch_invariant_mode=True, + flash_attention_version=_BIK_FA_VERSION, params_dtype=torch.bfloat16, normalization="RMSNorm", layernorm_epsilon=1e-5, @@ -154,6 +165,7 @@ def test_te_row_parallel_linear_batch_invariant_randomized(): hidden_dropout=0.0, attention_dropout=0.0, batch_invariant_mode=True, + flash_attention_version=_BIK_FA_VERSION, params_dtype=torch.bfloat16, normalization="RMSNorm", layernorm_epsilon=1e-5, @@ -200,6 +212,7 @@ def test_te_layernorm_column_parallel_linear_batch_invariant_randomized(): hidden_dropout=0.0, attention_dropout=0.0, batch_invariant_mode=True, + flash_attention_version=_BIK_FA_VERSION, params_dtype=torch.bfloat16, normalization="RMSNorm", layernorm_epsilon=1e-5, @@ -246,6 +259,7 @@ def test_te_norm_batch_invariant_randomized(): hidden_dropout=0.0, attention_dropout=0.0, batch_invariant_mode=True, + flash_attention_version=_BIK_FA_VERSION, params_dtype=torch.bfloat16, normalization="RMSNorm", layernorm_epsilon=1e-5, @@ -279,6 +293,7 @@ def test_column_parallel_linear_batch_invariant_randomized(): hidden_dropout=0.0, attention_dropout=0.0, batch_invariant_mode=True, + flash_attention_version=_BIK_FA_VERSION, params_dtype=torch.bfloat16, normalization="RMSNorm", layernorm_epsilon=1e-5, @@ -332,6 +347,7 @@ def test_te_attention_layer_batch_invariant_randomized(): hidden_dropout=0.0, attention_dropout=0.0, batch_invariant_mode=True, + flash_attention_version=_BIK_FA_VERSION, params_dtype=torch.bfloat16, normalization="RMSNorm", layernorm_epsilon=1e-5, @@ -422,6 +438,7 @@ def test_te_column_parallel_linear_parity(): hidden_dropout=0.0, attention_dropout=0.0, batch_invariant_mode=True, + flash_attention_version=_BIK_FA_VERSION, params_dtype=torch.bfloat16, normalization="RMSNorm", layernorm_epsilon=1e-5, @@ -517,6 +534,7 @@ def test_te_rmsnorm_parity(): hidden_dropout=0.0, attention_dropout=0.0, batch_invariant_mode=True, + flash_attention_version=_BIK_FA_VERSION, params_dtype=torch.bfloat16, normalization="RMSNorm", layernorm_epsilon=1e-5, @@ -596,6 +614,7 @@ def test_te_layernorm_linear_parity(): hidden_dropout=0.0, attention_dropout=0.0, batch_invariant_mode=True, + flash_attention_version=_BIK_FA_VERSION, params_dtype=torch.bfloat16, normalization="RMSNorm", layernorm_epsilon=1e-5, diff --git a/tests/unit_tests/transformer/test_transformer_block.py b/tests/unit_tests/transformer/test_transformer_block.py index 63add511bcd..0f19bc3dc95 100644 --- a/tests/unit_tests/transformer/test_transformer_block.py +++ b/tests/unit_tests/transformer/test_transformer_block.py @@ -14,6 +14,7 @@ from megatron.core.models.gpt.gpt_model import GPTModel from megatron.core.process_groups_config import ProcessGroupCollection from megatron.core.tensor_parallel.random import model_parallel_cuda_manual_seed +from megatron.core.transformer.attention import SelfAttention from megatron.core.transformer.enums import ModelType from megatron.core.transformer.pipeline_parallel_layer_layout import PipelineParallelLayerLayout from megatron.core.transformer.spec_utils import build_module @@ -77,13 +78,101 @@ def test_gpu_forward_full_checkpoint(self): def test_gpu_forward_full_checkpoint_fp8(self): self._run_full_checkpoint_test(fp8="e4m3") + def test_gpu_forward_full_checkpoint_dual_rope(self): + kv_channels = self.transformer_config.kv_channels + sequence_length = 32 + rotary_pos_emb = ( + torch.ones(sequence_length, 1, 1, kv_channels, device='cuda'), + torch.ones(sequence_length, 1, 1, kv_channels, device='cuda'), + ) + + def modify_arg_and_forward( + target_func, args, kwargs, target_name, target_index, arg_modifier_func + ): + args_list = list(args) + + if target_name in kwargs: + kwargs[target_name] = arg_modifier_func(kwargs[target_name]) + elif len(args_list) > target_index: + args_list[target_index] = arg_modifier_func(args_list[target_index]) + else: + raise RuntimeError( + f"Argument '{target_name}' at index {target_index} was not provided in args or kwargs." + ) + + return target_func(*args_list, **kwargs) + + class MockSelfAttentionWithDualRope(SelfAttention): + def forward(self, *args, **kwargs): + """Switch to either local or global RoPE embedding before forward.""" + + def arg_modifier_func(rotary_pos_emb): + assert isinstance(rotary_pos_emb, (tuple, list)) and len(rotary_pos_emb) == 2 + + if self.dual_rope_kind == "local_and_global": + assert rotary_pos_emb[0] is not None + assert rotary_pos_emb[1] is not None + + if self.layer_number % 2 == 0: + final_rotary_pos_emb = rotary_pos_emb[0] + else: + final_rotary_pos_emb = rotary_pos_emb[1] + elif self.dual_rope_kind == "local_only": + assert rotary_pos_emb[0] is not None + assert rotary_pos_emb[1] is None + + final_rotary_pos_emb = rotary_pos_emb[0] + elif self.dual_rope_kind == "global_only": + assert rotary_pos_emb[0] is None + assert rotary_pos_emb[1] is not None + + final_rotary_pos_emb = rotary_pos_emb[1] + else: + assert False, f"Unknown dual_rope_kind: {self.dual_rope_kind}" + + return final_rotary_pos_emb + + return modify_arg_and_forward( + super().forward, args, kwargs, "rotary_pos_emb", 5, arg_modifier_func + ) + + # Test non-Dual RoPE + self._run_full_checkpoint_test( + fp8=None, seq_len=sequence_length, rotary_pos_emb=rotary_pos_emb[0] + ) + + # Test Dual RoPE + self._run_full_checkpoint_test( + fp8=None, + seq_len=sequence_length, + attn_class=MockSelfAttentionWithDualRope, + rotary_pos_emb=rotary_pos_emb, + dual_rope_kind="local_and_global", + ) + self._run_full_checkpoint_test( + fp8=None, + seq_len=sequence_length, + attn_class=MockSelfAttentionWithDualRope, + rotary_pos_emb=(rotary_pos_emb[0], None), + dual_rope_kind="local_only", + ) + self._run_full_checkpoint_test( + fp8=None, + seq_len=sequence_length, + attn_class=MockSelfAttentionWithDualRope, + rotary_pos_emb=(None, rotary_pos_emb[1]), + dual_rope_kind="global_only", + ) + def test_gpu_forward_selective_checkpoint(self): self._run_selective_checkpoint_test(fp8=None) def test_gpu_forward_selective_checkpoint_fp8(self): self._run_selective_checkpoint_test(fp8="e4m3") - def _run_full_checkpoint_test(self, fp8): + def _run_full_checkpoint_test( + self, fp8, seq_len=None, attn_class=None, rotary_pos_emb=None, dual_rope_kind=None + ): transformer_config = self.transformer_config config = transformer_config config.recompute_granularity = 'full' @@ -93,11 +182,17 @@ def _run_full_checkpoint_test(self, fp8): full_transformer_block = TransformerBlock( config, get_gpt_layer_with_transformer_engine_spec() ) + if attn_class is not None: + for layer in full_transformer_block.layers: + layer.self_attention.__class__ = attn_class + assert not hasattr(layer.self_attention, "dual_rope_kind") + layer.self_attention.dual_rope_kind = dual_rope_kind + assert full_transformer_block.config.recompute_granularity == 'full' assert full_transformer_block.config.recompute_method == 'block' assert full_transformer_block.config.fp8 == fp8 - sequence_length = 32 + sequence_length = 32 if seq_len is None else seq_len micro_batch_size = 2 full_transformer_block.cuda() @@ -108,7 +203,9 @@ def _run_full_checkpoint_test(self, fp8): attention_mask = torch.ones((1, 1, sequence_length, sequence_length), dtype=bool).cuda() hidden_states = full_transformer_block( - hidden_states=hidden_states, attention_mask=attention_mask + hidden_states=hidden_states, + attention_mask=attention_mask, + rotary_pos_emb=rotary_pos_emb, ) assert hidden_states.shape[0] == sequence_length assert hidden_states.shape[1] == micro_batch_size diff --git a/tests/unit_tests/transformer/test_transformer_config.py b/tests/unit_tests/transformer/test_transformer_config.py new file mode 100644 index 00000000000..febb3842789 --- /dev/null +++ b/tests/unit_tests/transformer/test_transformer_config.py @@ -0,0 +1,32 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import pytest + +from megatron.core.transformer.transformer_config import TransformerConfig + + +def _make_overlap_config(mtp_num_layers: int | None) -> TransformerConfig: + return TransformerConfig( + num_layers=1, + hidden_size=128, + num_attention_heads=4, + num_moe_experts=2, + expert_model_parallel_size=2, + moe_token_dispatcher_type="alltoall", + overlap_moe_expert_parallel_comm=True, + bf16=True, + mtp_num_layers=mtp_num_layers, + ) + + +@pytest.mark.parametrize("mtp_num_layers", [None, 0, 1]) +def test_ep_a2a_overlap_accepts_supported_mtp_layer_counts(mtp_num_layers: int | None): + config = _make_overlap_config(mtp_num_layers) + + assert config.mtp_num_layers == mtp_num_layers + + +@pytest.mark.parametrize("mtp_num_layers", [-1, 2]) +def test_ep_a2a_overlap_rejects_unsupported_mtp_layer_counts(mtp_num_layers: int): + with pytest.raises(AssertionError, match="MTP supports at most one layer"): + _make_overlap_config(mtp_num_layers) diff --git a/tests/unit_tests/transformer/test_transformer_engine_grouped_linear.py b/tests/unit_tests/transformer/test_transformer_engine_grouped_linear.py index 467a6c31322..84ca109a15e 100644 --- a/tests/unit_tests/transformer/test_transformer_engine_grouped_linear.py +++ b/tests/unit_tests/transformer/test_transformer_engine_grouped_linear.py @@ -36,6 +36,56 @@ def _empty_load_args(): return {}, True, [], [], [] +@pytest.mark.parametrize(("parallel_mode", "partition_dim"), (("column", 0), ("row", 1))) +def test_expert_parameter_attributes_use_expert_topology(parallel_mode, partition_dim): + module = torch.nn.Module() + module.register_parameter("weight0", torch.nn.Parameter(torch.empty(4, 4))) + module.register_parameter("bias0", torch.nn.Parameter(torch.empty(4))) + + te_ext._set_expert_parameter_attributes( + module, parallel_mode=parallel_mode, use_expert_pgs=True + ) + + assert module.weight0.allreduce is False + assert module.weight0.tensor_model_parallel is True + assert module.weight0.partition_dim == partition_dim + assert module.bias0.allreduce is False + assert module.bias0.tensor_model_parallel is (parallel_mode == "column") + + +@pytest.mark.parametrize( + ("name", "is_partitioned"), + ( + ("weight", True), + ("weight12", True), + ("bias", True), + ("bias12", True), + ("weight_scale", False), + ("bias_extra", False), + ), +) +def test_expert_parameter_attributes_match_parameter_names(name, is_partitioned): + module = torch.nn.Module() + module.register_parameter(name, torch.nn.Parameter(torch.empty(4))) + + te_ext._set_expert_parameter_attributes(module, parallel_mode="column", use_expert_pgs=True) + + param = module.get_parameter(name) + assert getattr(param, "tensor_model_parallel", False) is is_partitioned + + +def test_split_empty_extra_state_for_stateless_recipe(): + module = _grouped_linear_stub(num_gemms=2) + module.fp8_meta = {"fp8_checkpoint": True} + module.fp8 = False + module.fp8_calibration = False + + states = module._split_extra_state(torch.empty(0, dtype=torch.uint8)) + + assert len(states) == 2 + assert all(state.dtype == torch.uint8 and state.numel() == 0 for state in states) + + def test_split_grouped_checkpoint_tensor_uses_quantized_members(): module = _grouped_linear_stub(num_gemms=2) members = [torch.tensor([1, 2]), torch.tensor([3, 4])] diff --git a/tools/check_golden_values.py b/tools/check_golden_values.py new file mode 100644 index 00000000000..270786f496f --- /dev/null +++ b/tools/check_golden_values.py @@ -0,0 +1,82 @@ +# Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +"""Check golden-value JSON files for NaN and infinity values.""" + +import argparse +import json +import logging +from collections.abc import Iterator +from pathlib import Path +from typing import Any + +logger = logging.getLogger(__name__) + +NOT_ACCEPTED_VALUES = [ + "nan", + "+nan", + "-nan", + "inf", + "+inf", + "-inf", + "infinity", + "+infinity", + "-infinity", +] + + +def _find_non_finite_values(value: Any, location: str = "$") -> Iterator[tuple[str, Any]]: + if isinstance(value, dict): + for key, child in value.items(): + yield from _find_non_finite_values(child, f"{location}[{key!r}]") + elif isinstance(value, list): + for index, child in enumerate(value): + yield from _find_non_finite_values(child, f"{location}[{index}]") + elif str(value).strip().lower() in NOT_ACCEPTED_VALUES: + yield location, value + + +def _format_failures(failures: list[tuple[str, Any]], limit: int = 20) -> str: + lines = [f" {location} = {value!r}" for location, value in failures[:limit]] + if len(failures) > limit: + lines.append(f" ... and {len(failures) - limit} more") + return "\n".join(lines) + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Fail if any golden-value JSON file contains NaN or infinity." + ) + parser.add_argument("files", nargs="+", type=Path, help="Golden-value JSON files to check.") + return parser.parse_args() + + +def main() -> int: + """Check the requested golden-value files and return a process exit code.""" + failed = False + files = _parse_args().files + + for golden_value_file in files: + try: + with golden_value_file.open() as file: + golden_values = json.load(file) + except (OSError, json.JSONDecodeError) as error: + logger.error("Could not read %s: %s", golden_value_file, error) + failed = True + continue + + failures = list(_find_non_finite_values(golden_values)) + if failures: + logger.error( + "Found non-finite values in %s:\n%s", golden_value_file, _format_failures(failures) + ) + failed = True + + if not failed: + logger.info("Checked %d golden-value file(s); all values are finite.", len(files)) + + return int(failed) + + +if __name__ == "__main__": + logging.basicConfig(level=logging.INFO, format="%(message)s") + raise SystemExit(main()) diff --git a/tools/common_pile_dataset/README.md b/tools/common_pile_dataset/README.md index 2431d1b01d3..dcc18fee24b 100644 --- a/tools/common_pile_dataset/README.md +++ b/tools/common_pile_dataset/README.md @@ -165,7 +165,7 @@ default. On HPC systems where `/home` is small, set `HF_HOME` to a path with sufficient space: ```bash -export HF_HOME=/lustre/path/to/.hf_cache +export HF_HOME=/lustre/path/to/hf_home ``` The setup script does this automatically. diff --git a/tools/common_pile_dataset/setup_common_pile_dataset.sh b/tools/common_pile_dataset/setup_common_pile_dataset.sh index cb869e28368..01c438cd21a 100644 --- a/tools/common_pile_dataset/setup_common_pile_dataset.sh +++ b/tools/common_pile_dataset/setup_common_pile_dataset.sh @@ -29,7 +29,7 @@ DATASET_NAME="common-pile/comma_v0.1_training_dataset" WORK_DIR="/tmp/mcore_dataset_setup_$$" # Redirect HuggingFace cache to lustre so it doesn't fill up /home -export HF_HOME="/lustre/fsw/portfolios/coreai/projects/coreai_dlalgo_mcore/mcore_ci/.hf_cache" +export HF_HOME="/lustre/fsw/portfolios/coreai/projects/coreai_dlalgo_mcore/mcore_ci/hf_home" export HF_DATASETS_CACHE="${HF_HOME}/datasets" echo "============================================================"