Skip to content

[ROCm] support warpSize 32 and 64 in a single TBE codegen build - #6123

Open
jeffdaily wants to merge 1 commit into
pytorch:mainfrom
jeffdaily:warpsize-5-tbe
Open

jeffdaily wants to merge 1 commit into
pytorch:mainfrom
jeffdaily:warpsize-5-tbe

Conversation

@jeffdaily

@jeffdaily jeffdaily commented Aug 7, 2026 •

Copy link
Copy Markdown
Contributor

Make the TBE codegen emit kernels and host dispatch tables that are correct for both wave32 and wave64 AMD archs from a single build, including a mixed PYTORCH_ROCM_ARCH="gfx90a;gfx1100" wheel. This is the codegen core of the warpSize port; the earlier PRs in the chain supplied the primitives it needs (kWarpSizeHost, warpReduceAllSum gating, cache associativity, and the CMake-derived --has_wave32/--has_wave64 flags). The mechanical host-side kWarpSize -> kWarpSizeHost() launch-config substitutions were split into #6278 at reviewer request; this PR carries the codegen core and is stacked on #6278.

Two root causes are fixed here:

  1. Kernel-name mangling baked in the wave size. The TBE kernel templates took kThreadGroupSize as a template integer, so its value entered the mangled symbol name; on a mixed-arch build the host pass and the gfx1100 device pass disagree and the linker cannot resolve the device symbol. The templates are reparameterised to take kSubwarpDivisor (a wave-size-free literal); the kernel body recovers kThreadGroupSize = kWarpSize / kSubwarpDivisor in the device pass.
  2. Host dispatch used the wrong wave size. The DISPATCH_* macros emit a _WAVE32/_WAVE64 pair selected at runtime via kWarpSizeHost(); single-wave-size builds emit only the matching table.

Reconciliation with main's HIP-optimized backward path (#5712, #6197): the hip_mixed_d warp kernel gets the same kSubwarpDivisor treatment as the other kernels (definition, explicit instantiations, and host selection). Its max_D <= 128 fast path re-instantiates with a 32-lane thread group, which under this parameterisation is kSubwarpDivisor = 2 on wave64; like the matching cta_per_row override, it is emitted only for wave64 builds and guarded by a runtime kWarpSizeHost() check on mixed wheels, since on wave32 the default full-warp kernel already runs 32-lane groups.

Two related wave32 device-pass fixes are folded in (num_packed_bags cap in nbit inference; runtime max_vecs_per_thread in the backward template for dims >= 128).

Fifth and final in the chain splitting #5804 into reviewable pieces; stacked on #6278.

Authored with assistance from Claude (Anthropic).

@meta-codesync

meta-codesync Bot commented Aug 7, 2026

Copy link
Copy Markdown
Contributor

@q10 has imported this pull request. If you are a Meta employee, you can view this in D115263090.

@jeffdaily

Copy link
Copy Markdown
Contributor Author

@q10 Friendly ping on this one -- it's been imported (D115263090) for about a week and is green/mergeable. This is the last of the five-PR split of #5804; it's ready to land whenever the internal diff can go through. Thanks!

Authored with assistance from Claude (Anthropic).

@q10

q10 commented Aug 12, 2026

Copy link
Copy Markdown
Contributor

@q10 Friendly ping on this one -- it's been imported (D115263090) for about a week and is green/mergeable. This is the last of the five-PR split of #5804; it's ready to land whenever the internal diff can go through. Thanks!

Authored with assistance from Claude (Anthropic).

@jeffdaily yes, we are looking into this

utils::cuda::cap_grid_dim_x(
nbit::div_round_up(
B * T + 1, kForwardMaxThreads / kWarpSize),
B * T + 1, kForwardMaxThreads / kWarpSizeHost()),

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we split the change from kWarpSize to kWarpSizeHost() into a separate PR, in order to reduce the impact of each PR and make rollback easier if there is ever a need for rollback?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done. The kWarpSize -> kWarpSizeHost() host-launch substitutions (including the hunk on this line) are split out into #6278, and this PR is rebased on top of it so it now carries only the codegen/template reparameterisation. While rebasing over main I also reconciled with the hip_mixed_d warp kernel from #5712; the updated PR description has the details.

Authored with assistance from Claude (Anthropic).

@pytorch-bot

pytorch-bot Bot commented Sep 9, 2026 •

Copy link
Copy Markdown

Workflows were awaiting approval. CI has now been triggered for the ciflow labels on this PR.

jeffdaily added a commit to jeffdaily/FBGEMM that referenced this pull request Sep 15, 2026
Host-side TBE launch configurations computed block/grid dims from the
compile-time kWarpSize (or a hardcoded 64 for BT_block_size). In a ROCm wheel
that serves both wave64 (CDNA) and wave32 (RDNA) archs, the host code is
compiled once, so a compile-time warp size is wrong for whichever arch it does
not match: on gfx1100 the launches came up with 64-wide x-dims for 32-wide
warps. This replaces those host uses with the runtime kWarpSizeHost(), which
reports the active device's warp size. On CUDA and on wave64-only ROCm builds
kWarpSizeHost() folds to the same values as before, so this is
behavior-preserving there.

Split out of pytorch#6123 at reviewer request to keep each PR small and
independently revertible. The codegen/template reparameterisation that makes
the kernel symbols themselves wave-size-agnostic stays in pytorch#6123, which is
rebased on top of this.

Test Plan:

Rendered the TBE codegen for CUDA, ROCm wave64, ROCm wave32, and ROCm
wave32+wave64 from this commit and from origin/main, and diffed: the only
changes are the intended kWarpSize -> kWarpSizeHost() substitutions in host
launch code. Also compiled fbgemm_gpu for gfx90a with the follow-up commit
applied (see pytorch#6123).

    cd fbgemm_gpu/codegen/genscript
    for flags in "" "--is_rocm --has_wave64" "--is_rocm --has_wave32" \
                 "--is_rocm --has_wave32 --has_wave64"; do
      python generate_backward_split.py --opensource $flags
      python generate_forward_split.py --opensource $flags
      python generate_forward_quantized.py --opensource $flags
      python generate_embedding_optimizer.py --opensource $flags
      python generate_index_select.py --opensource $flags
    done

Authored with assistance from Claude (Anthropic).
meta-codesync Bot pushed a commit that referenced this pull request Sep 16, 2026
Summary:
X-link: https://github.com/facebookresearch/FBGEMM/pull/3171

Host-side TBE launch configurations computed block/grid dims from the compile-time `kWarpSize` (or a hardcoded 64 for `BT_block_size`). In a ROCm wheel that serves both wave64 (CDNA) and wave32 (RDNA) architectures, the host code is compiled once, so a compile-time warp size is wrong for whichever architecture it does not match.

This replaces those host uses with runtime warp-size queries: `kWarpSizeHost()` in generated CUDA sources and a guarded `at::cuda::warp_size()` call in the host-generated backward code. The latter preserves the prior wave64 fallback for CPU and Meta dispatch and avoids pulling device-only declarations into a host compilation unit. Launch sites cache the runtime value so each configuration uses one consistent geometry.

On CUDA and wave64-only ROCm builds, the runtime query returns the same value as before. This was split out of #6123 at reviewer request to keep each PR small and independently revertible. The codegen/template reparameterisation that makes kernel symbols wave-size-agnostic stays in that follow-up.

Authored with assistance from Claude (Anthropic).

Pull Request resolved: #6278

Test Plan:
# Script

```
arc f fbcode/deeplearning/fbgemm/fbgemm_gpu/codegen/inference/embedding_forward_quantized_split_lookup.cu fbcode/deeplearning/fbgemm/fbgemm_gpu/codegen/inference/embedding_forward_quantized_split_nbit_host_template.cu fbcode/deeplearning/fbgemm/fbgemm_gpu/codegen/training/backward/embedding_backward_split_host_template.cpp fbcode/deeplearning/fbgemm/fbgemm_gpu/codegen/training/backward/embedding_backward_split_indice_weights_template.cu fbcode/deeplearning/fbgemm/fbgemm_gpu/codegen/training/forward/embedding_forward_split_template.cu
arc lint -e extra fbcode/deeplearning/fbgemm/fbgemm_gpu/codegen/inference/embedding_forward_quantized_split_lookup.cu fbcode/deeplearning/fbgemm/fbgemm_gpu/codegen/inference/embedding_forward_quantized_split_nbit_host_template.cu fbcode/deeplearning/fbgemm/fbgemm_gpu/codegen/training/backward/embedding_backward_split_host_template.cpp fbcode/deeplearning/fbgemm/fbgemm_gpu/codegen/training/backward/embedding_backward_split_indice_weights_template.cu fbcode/deeplearning/fbgemm/fbgemm_gpu/codegen/training/forward/embedding_forward_split_template.cu
buck build --flagfile fbcode//mode/opt fbcode//deeplearning/fbgemm/fbgemm_gpu/codegen:embedding_ops_inference_gpu fbcode//deeplearning/fbgemm/fbgemm_gpu/codegen:embedding_ops_training_gpu
```

# Results

All formatting, lint, and focused Buck build commands passed.

# Analysis of results

Formatting and lint checks passed. The focused Buck build compiled both generated inference and training owners, including the updated dispatch-device predicate and the nobag-only launch geometry.

# Other test plan info from agent

Resubmission triggers CI. No local gfx1100/gfx90a runtime hardware test was available.

Reviewed By: yvonne-lab

Differential Revision: D119363782

Pulled By: q10

fbshipit-source-id: 3e731c4287fa9549151ef57fb7e6eb6425294620
@q10

q10 commented Sep 17, 2026

Copy link
Copy Markdown
Contributor

@jeffdaily could you rebase this PR on top of latest main to address the merge conflicts?

Make the TBE codegen emit kernels and host dispatch tables that are correct
for both wave32 and wave64 AMD archs from a single build, including a mixed
PYTORCH_ROCM_ARCH="gfx90a;gfx1100" wheel. This is the codegen core of the
warpSize port; the earlier PRs in the chain supplied the primitives it needs
(kWarpSizeHost, warpReduceAllSum gating, cache associativity, and the
CMake-derived --has_wave32/--has_wave64 flags). The mechanical host-side
kWarpSize -> kWarpSizeHost() launch-config substitutions were split into the
preceding commit at reviewer request; this commit carries the codegen core.

Two root causes are fixed here.

1. Kernel-name mangling baked in the wave size. The TBE kernel templates took
   kThreadGroupSize as a template integer, so its value (derived from warpSize)
   entered the mangled symbol name. On a mixed-arch build the host pass
   instantiates with the ROCm-64 placeholder while the gfx1100 device pass
   instantiates with 32, and the linker cannot resolve the gfx1100 device
   symbol. The templates are reparameterised to take kSubwarpDivisor (a
   wave-size-free literal); the kernel body recovers
   kThreadGroupSize = kWarpSize / kSubwarpDivisor in the device pass. Host and
   every device pass now agree on the symbol, and one instantiation per bracket
   serves both waves (~30% fewer symbols than emitting wave32 and wave64
   variants separately).

2. Host dispatch used the wrong wave size. The DISPATCH_* macros and launch
   configs are reworked to emit a _WAVE32 / _WAVE64 pair and select at runtime
   via kWarpSizeHost(); single-wave-size builds (driven by has_wave32 /
   has_wave64) emit only the matching table, so single-arch wheels do not grow.
   jinja_environment.py gains get_max_vecs_template_configs_union[_forward] to
   emit the union of instantiations needed by the enabled wave sizes.

Two related fixes for the wave32 device pass are folded in because they live
in these same templates: capping num_packed_bags by warp lane capacity in nbit
inference, and deriving max_vecs_per_thread from the runtime warp width in the
backward template (a compile-time items_per_warp baked to 256 on ROCm dropped
gradients for embedding dims >= 128 on wave32).

Reconciliation with main's HIP-optimized backward path (pytorch#5712, pytorch#6197): the
hip_mixed_d warp kernel gets the same kSubwarpDivisor treatment as the other
kernels (definition, explicit instantiations, and host selection). Its
max_D <= 128 fast path re-instantiates with a 32-lane thread group, which
under this parameterisation is kSubwarpDivisor = 2 on wave64; like the
matching cta_per_row override, it is emitted only for wave64 builds and
guarded by a runtime kWarpSizeHost() check on mixed wheels, since on wave32
the default full-warp kernel already runs 32-lane groups.

Suggested review order: jinja_environment.py (the codegen helpers and dispatch
macros) first, then the forward/backward/optimizer templates that consume them,
then the inference templates.

Fifth and final in the chain splitting pytorch#5804 into reviewable
pieces; stacked on the kWarpSizeHost() launch-config commit.

Test Plan:

Rendered the TBE codegen for CUDA, ROCm wave64, ROCm wave32, and ROCm
wave32+wave64 and compared against origin/main renders: CUDA and wave64
instantiation sets match main one-for-one (group sizes map to divisors), the
max_D <= 128 overrides appear unconditionally on wave64-only builds, behind a
runtime kWarpSizeHost() == 64 check on mixed builds, and not at all on wave32
or CUDA builds. Built fbgemm_gpu for gfx90a (wave64) from this commit and ran
the TBE training suites on gfx90a: forward 13 passed / 4 skipped,
backward_adagrad 12 passed / 5 skipped, backward sgd+dense+none 7 passed /
5 skipped, no failures.

    cd fbgemm_gpu
    export PYTORCH_ROCM_ARCH=gfx90a BUILD_ROCM_VERSION=7.14
    python setup.py build --build-variant=rocm --build-target=default
    cd test/tbe/training
    python -m pytest forward_test.py -q
    python -m pytest backward_adagrad_test.py -q
    python -m pytest backward_sgd_test.py backward_dense_test.py backward_none_test.py -q

The wave32 and mixed-arch (gfx90a;gfx1100) paths -- link resolution and
numerics -- were validated on gfx1100 in pytorch#5804.

Authored with assistance from Claude (Anthropic).
@jeffdaily

Copy link
Copy Markdown
Contributor Author

@q10 rebased onto latest main. The conflicts are gone and this is now a single commit, since the two commits below it landed with #6278.

The rebase needed one reconciliation worth a look: #6286 added a ROCm half-wave-group fast path to the forward dispatch macro using the old kThreadGroupSize parameterization, which this PR replaces with kSubwarpDivisor. It is now in the _WAVE64 macro as kSubwarpDivisor = 2, with its explicit instantiation gated on has_wave64. Wave32 does not take the override because its default MAX_D <= 128 dispatch is already a 32-lane group.

Checked by rendering the codegen for CUDA, wave64, wave32 and mixed, and comparing kernel instantiation multisets against main per file and per kernel with main's group size G mapped to this PR's divisor D via G = warpSize / D. CUDA and wave64 match one for one: 77644 and 90325 instantiations on both sides.

CI here is in action_required as well.

Authored with assistance from Claude (Anthropic).

@jeffdaily

Copy link
Copy Markdown
Contributor Author

@q10 Checking in on this one now that the rest of the chain is in. #6326 landed the ROCm backward launch-config restore, so this PR is the last piece of the #5804 split. It has been rebased onto main since 9/18, it's a single commit with no conflicts, and the ROCm and CUDA CI jobs are green.

The only import is still D115263090 from August, which predates the rebase. Could you re-import it against the current head? The two "Facebook Internal" checks (Builds & Tests, Linter) failed on 9/18 and we can't see those results from outside. If they point to a real problem, let me know what they report and I'll fix it.

Authored with assistance from Claude (Anthropic).

@q10

q10 commented Sep 29, 2026

Copy link
Copy Markdown
Contributor

Hi @jeffdaily thanks for checking in. We have been working to verify this PR. I will import the latest changes from this PR and re-test against the stack of code changes we have. We hope to land this soon as it will unblock us for MI450 developments.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants