This is an automated email from the ASF dual-hosted git repository.

spectrometerHBH pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git


The following commit(s) were added to refs/heads/main by this push:
     new 771193a5c2 [FEATURE][TIRx][CUDA] Support TMEM datapath B (#20075)
771193a5c2 is described below

commit 771193a5c26188713b2ec01f5bcd937bf7966fae
Author: Hongyi Jin <[email protected]>
AuthorDate: Thu Jul 30 20:40:18 2026 -0400

    [FEATURE][TIRx][CUDA] Support TMEM datapath B (#20075)
    
    ## Motivation and context
    
    `tcgen05.mma.cta_group::2` describes an M=128 operation across a pair of
    CTAs, but each CTA owns a logical `(64, N)` accumulator tile. That
    per-CTA tile uses TMEM datapath B: its low `N/2` columns occupy physical
    lanes 0–63 and its high `N/2` columns occupy lanes 64–127:
    
    | Logical coordinate | Physical TMEM coordinate |
    | --- | --- |
    | `(r, c)` | `(TLane = r + 64 * floor(c / (N/2)), TCol = c mod (N/2))` |
    
    The resulting logical `(64, N)` tile has the same physical footprint as
    a conventional `(128, N/2)` Layout D tile.
    
    Before this PR, the `cta_group=2` GEMM lowering constructed this mapping
    locally for the accumulator write, but the public TMEM allocation and
    `tcgen05.ld/st` paths understood only datapaths D and F. Code could
    produce a datapath-B accumulator, but could not describe, validate, and
    read or write it through the normal TIRx layout and fragment APIs.
    Treating it as an ordinary M=64 `.16x*b` fragment accesses the wrong
    lanes, while exposing a `(128, N/2)` alias loses the logical per-CTA
    shape.
    
    This PR makes datapath B a first-class layout shared by allocation, GEMM
    lowering, and register transfer. A logical `(64, N)` fp32 fragment is
    transferred as the physical `.32x32b.x{N/2}` image over all 128 lanes.
    
    ## Changes
    
    - Add `tmem_datapath_layout("B", 64, N)` and `tmem_pool.alloc(...,
    datapath="B")`.
    - Single-source the `cta_group=2` GEMM accumulator mapping through that
    public layout factory.
    - Support `T.alloc_tcgen05_ldst_frag("32x32b", (64, N), "float32")` as
    the register image of datapath B.
    - Classify datapath B before the ordinary D/F atom paths and emit one
    physical `tcgen05.{ld,st}.32x32b.x{N/2}` operation.
    - Require an even `N`, a PTX-supported `N/2`, fp32 cells, matching
    layouts, and a full logical `(64, N)` transfer. Partial logical-column
    slices are rejected because they are not physically contiguous after the
    lane split.
    - Document the logical-to-physical mapping, supported APIs, and transfer
    restrictions.
    - Update existing `cta_group=2` expected IR to retain the `shared`
    storage scope required by the pointer-type fix that is now on `main`.
    
    ## Testing
    
    - Directly verify Layout B and register-fragment `(row, col) -> (lane,
    register)` mappings.
    - Verify invalid shapes, dtypes, fragments, sub-slab values, and partial
    transfers are rejected.
    - Check both load and store codegen for the expected `.32x32b`
    instruction.
    - On B200, round-trip every register through TMEM for `N = 32, 64, 128,
    256` at base-column offsets 0 and 32.
    - On B200, run a `cta_group=2` GEMM and read the Layout B accumulator
    back through the public fragment API.
    - `194 passed` across:
    -
    
`tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tmem_16xnb.py`
    -
    
`tests/python/tirx/operator/tile_primitive/cuda/gemm_async/test_gemm_async.py`
    - Changed-files pre-commit checks pass.
    
    ## Dependencies
    
    The tcgen05 fragment-layout foundation from #20068 is already merged.
    This PR is based directly on current `apache/main`.
---
 docs/tirx/layout.rst                               |  27 +++
 docs/tirx/native_basics/cuda/buffers.rst           |   6 +-
 .../tile_primitives/copy_async/tcgen05_ldst.rst    |  48 ++++-
 docs/tirx/tile_primitives/gemm_async.rst           |  42 ++++-
 python/tvm/backend/cuda/lang/alloc_pool.py         |   3 +-
 .../tile_primitive/copy_async/tcgen05_ldst.py      | 109 ++++++++++-
 .../operator/tile_primitive/gemm_async/tcgen05.py  |   6 +-
 python/tvm/tirx/layout.py                          | 114 +++++++++---
 python/tvm/tirx/script/builder/ir.py               |  14 +-
 .../cuda/copy_async/test_tmem_16xnb.py             | 204 +++++++++++++++++++++
 .../cuda/gemm_async/test_gemm_async.py             | 145 ++++++++++++++-
 11 files changed, 664 insertions(+), 54 deletions(-)

diff --git a/docs/tirx/layout.rst b/docs/tirx/layout.rst
index d8bff61f62..a71bb9f319 100644
--- a/docs/tirx/layout.rst
+++ b/docs/tirx/layout.rst
@@ -269,6 +269,7 @@ TMEM datapath layouts
     accum = tmem_datapath_layout("D", 128, cols)
     lower = tmem_datapath_layout("F", 64, cols, sub_slab=0)
     upper = tmem_datapath_layout("F", 64, cols, sub_slab=1)
+    paired = tmem_datapath_layout("B", 64, cols)
 
 Layout D maps logical row ``r`` directly to ``TLane = r`` and spans both
 16-lane halves of every warp's 32-lane TMEM partition. Layout F maps its 64
@@ -287,6 +288,32 @@ upper-half aliases of the same 128-row Layout D 
allocation. The
 ``row=0`` or ``row=16`` instruction. Layout D already occupies both halves,
 so a nonzero ``sub_slab`` is rejected.
 
+Layout B is the per-CTA accumulator placement for an M=64
+``tcgen05.mma.cta_group::2`` operation. PTX describes the two-CTA operation
+as M=128, while each CTA owns a logical ``(64, N)`` tile. Its columns split
+across the two 64-lane halves:
+
+.. math::
+
+   \mathrm{TLane}
+   = r + 64\left\lfloor\frac{c}{N/2}\right\rfloor,
+   \qquad
+   \mathrm{TCol} = c \bmod (N/2).
+
+Thus the tile occupies all 128 lanes and ``N/2`` tensor-memory columns, the
+same physical footprint as Layout D ``(128, N/2)``. ``N`` must be even and
+Layout B does not accept ``sub_slab``. Its register image uses the existing
+fragment API:
+
+.. code-block:: python
+
+    frag = T.alloc_tcgen05_ldst_frag("32x32b", (64, N), "float32")
+    Tx.wg.copy_async(frag[:, :], paired_accumulator[:, :])
+    T.ptx.tcgen05.wait.ld()
+
+The logical ``(64, N)`` fragment is one physical ``.32x32b`` transfer over
+all 128 lanes; each thread owns ``N/2`` contiguous fp32 registers.
+
 Beyond GPU registers
 ~~~~~~~~~~~~~~~~~~~~~~
 
diff --git a/docs/tirx/native_basics/cuda/buffers.rst 
b/docs/tirx/native_basics/cuda/buffers.rst
index 18de83ea03..9977725ead 100644
--- a/docs/tirx/native_basics/cuda/buffers.rst
+++ b/docs/tirx/native_basics/cuda/buffers.rst
@@ -379,7 +379,7 @@ tensor as a view at a column offset, and one warp frees it 
at the end:
         T.ptx.tcgen05.relinquish_alloc_permit(cta_group=cta_group)
         T.ptx.tcgen05.dealloc(addr, n_cols=512, cta_group=cta_group)
 
-You manage the column offsets and the ``tmem_layout`` (a datapath D/F layout)
+You manage the column offsets and the ``tmem_layout`` (a datapath D/F/B layout)
 yourself. This is exactly the sequence the pool below emits.
 
 Pool
@@ -393,7 +393,9 @@ bump-allocation, and the datapath layout:
     tmem_addr = pool.alloc((1,), "uint32")          # pool = the kernel's smem 
pool
     tmem_pool = T.TMEMPool(pool, total_cols=512, cta_group=cta_group,
                            tmem_addr=tmem_addr)
-    acc = tmem_pool.alloc((CTA_M, 512), "float32")  # allocated_addr set for 
you
+    # Choose the layout required by the instruction that consumes the buffer:
+    acc = tmem_pool.alloc((CTA_M, 512), "float32")  # Layout D when CTA_M=128
+    # acc = tmem_pool.alloc((64, N), "float32", datapath="B")  # cta_group=2
     tmem_pool.commit()                               # emits tcgen05.alloc 
(one warp)
     # ... use acc ...
     tmem_pool.dealloc()                              # emits tcgen05.dealloc 
(one warp)
diff --git a/docs/tirx/tile_primitives/copy_async/tcgen05_ldst.rst 
b/docs/tirx/tile_primitives/copy_async/tcgen05_ldst.rst
index d34147f60e..106e084445 100644
--- a/docs/tirx/tile_primitives/copy_async/tcgen05_ldst.rst
+++ b/docs/tirx/tile_primitives/copy_async/tcgen05_ldst.rst
@@ -59,10 +59,12 @@ lowering:
      - ``(tmem, local)`` or ``(local, tmem)`` — exactly one side is tensor 
memory
    * - register layout
      - matched against a ``tcgen05_atom_layout`` (``.16x64b`` / ``.16x128b`` /
-       ``.16x256b``) for the fast path; otherwise the ``.32x32b`` fallback
+       ``.16x256b``) for the fast path; otherwise the ``.32x32b`` fallback.
+       Layout B uses the special ``.32x32b`` logical ``(64, N)`` image
    * - tmem datapath
-     - classified ``D`` (M=128 identity) or ``F`` (M=64 scattered); an F layout
-       also selects the lower or upper 16-lane sub-slab of each warp partition
+     - classified ``D`` (M=128 identity), ``F`` (M=64 scattered), or ``B``
+       (per-CTA M=64, ``cta_group=2`` column split); an F layout also selects
+       the lower or upper 16-lane sub-slab of each warp partition
 
 Demonstration program
 ----------------------
@@ -128,6 +130,10 @@ fragment spans two 16-row slabs, so the warps issue the 
atom twice
            shape=shape, num=num_eff,
            row=(sub_slab + slab) * 16, col=col_off_32b)
 
+Layout B is routed before the ordinary atom matching. Its logical
+``(64, N)`` fragment is emitted as one physical ``.32x32b.x{N/2}``
+instruction over all 128 lanes.
+
 The dispatch emits **no** wait — the caller issues ``tcgen05.wait.ld()`` /
 ``wait.st()`` (as in the demo).
 
@@ -166,6 +172,36 @@ The lower view emits ``row=0`` and the upper view emits 
``row=16`` for
 ``.16x64b``, ``.16x128b``, and ``.16x256b`` atoms. Layout D has 128 rows and
 already spans both sub-slabs, so it only accepts ``sub_slab=0``.
 
+Layout B readback
+-----------------
+
+Layout B is produced by an M=64-per-CTA ``tcgen05.mma`` with
+``cta_group=2``. Its logical ``(64, N)`` columns are split into two
+``N/2`` halves: the low half uses physical TMEM lanes 0–63, and the high
+half uses lanes 64–127. It therefore occupies all 128 lanes and ``N/2``
+``TCol`` values.
+
+Use the public allocation and fragment APIs together:
+
+.. code-block:: python
+
+    accumulator = tmem_pool.alloc((64, N), "float32", datapath="B")
+    frag = T.alloc_tcgen05_ldst_frag("32x32b", (64, N), "float32")
+
+    Tx.wg.copy_async(frag[:, :], accumulator[:, :])
+    T.ptx.tcgen05.wait.ld()
+
+    # The inverse direction emits tcgen05.st with the same physical image.
+    Tx.wg.copy_async(accumulator[:, :], frag[:, :])
+    T.ptx.tcgen05.wait.st()
+
+This is a single ``tcgen05.{ld,st}.32x32b.x{N/2}`` issue. ``N`` must be
+even, ``N/2`` must be a valid PTX ``num``, and the fragment is fp32-only.
+The first implementation intentionally requires the full logical
+``(64, N)`` region: a partial logical-column slice is not contiguous after
+the two-way lane split. A conventional ``.16x*b`` fragment is rejected with
+an actionable error.
+
 Generated TIRx IR
 -----------------
 
@@ -207,7 +243,9 @@ How inputs change the algorithm
    * - direction
      - ``tmem → local`` → ``tcgen05.ld``; ``local → tmem`` → ``tcgen05.st`` 
(same
        shape/num logic)
-   * - datapath D vs F
+   * - datapath D vs F vs B
      - ``D`` (M=128) covers all 128 rows; an M=128 ``.16x*b`` copy issues two 
slabs
        (``row = 0`` / ``row = 16``). ``F`` (M=64) scatters rows to lanes and
-       its layout selects one issue at ``row = 0`` or ``row = 16``
+       its layout selects one issue at ``row = 0`` or ``row = 16``. ``B``
+       (per-CTA M=64, ``cta_group=2``) splits N across two 64-lane halves and
+       emits one logical-64-row ``.32x32b`` image
diff --git a/docs/tirx/tile_primitives/gemm_async.rst 
b/docs/tirx/tile_primitives/gemm_async.rst
index 1f784f5c2b..c25cdfd57b 100644
--- a/docs/tirx/tile_primitives/gemm_async.rst
+++ b/docs/tirx/tile_primitives/gemm_async.rst
@@ -57,9 +57,9 @@ A single predicate — single-thread or warp scope:
        ``float8_e4m3fn`` / ``float4_e2m1fn`` with ``SFA`` / ``SFB`` scale 
factors in
        tmem; accumulator always ``float32``
    * - shape
-     - ``M ∈ {64, 128}`` (×2 for cta_group=2); ``N`` divisible by 8 
(cta_group=1) or
-       16 (cta_group=2); ``K`` divisible by ``MMA_K`` = 16 (f16/bf16) / 32 
(fp8) /
-       64 (fp4)
+     - per-CTA ``M ∈ {64, 128}``; ``N`` divisible by 8 (cta_group=1) or 16
+       (cta_group=2); ``K`` divisible by ``MMA_K`` = 16 (f16/bf16) / 32 (fp8) /
+       64 (fp4). With cta_group=2, the CTA pair covers twice the per-CTA M
    * - cta_group
      - ``1`` (one CTA) or ``2`` (two CTAs split the operand)
 
@@ -126,6 +126,39 @@ instruction descriptor is encoded at runtime. As with the 
other async ops, the
 dispatch emits **no** completion — the caller's ``tcgen05.commit`` + mbarrier 
wait
 close it.
 
+Accumulator datapaths and readback
+----------------------------------
+
+The accumulator layout must match the MMA's row placement:
+
+* Layout D is the M=128 identity placement.
+* Layout F is the single-CTA M=64 scattered placement.
+* Layout B is the per-CTA M=64 placement for ``cta_group=2``. Its logical
+  N columns split across physical lane halves 0–63 and 64–127, so it
+  occupies all 128 lanes and ``N/2`` tensor-memory columns.
+
+Allocate and read a Layout B result as follows:
+
+.. code-block:: python
+
+    accumulator = tmem_pool.alloc((64, N), "float32", datapath="B")
+    Tx.gemm_async(
+        accumulator[:, :],
+        A_smem[:, :],
+        B_smem[:, :],
+        dispatch="tcgen05",
+        cta_group=2,
+    )
+
+    frag = T.alloc_tcgen05_ldst_frag("32x32b", (64, N), "float32")
+    Tx.wg.copy_async(frag[:, :], accumulator[:, :])
+    T.ptx.tcgen05.wait.ld()
+
+The fragment is a logical ``(64, N)`` view of one physical
+``.32x32b`` transfer over all 128 lanes. The gemm write-side layout and
+the allocation/readback layout are produced by the same
+``tmem_datapath_layout("B", 64, N)`` factory.
+
 Generated TIRx IR
 -----------------
 
@@ -167,7 +200,8 @@ How inputs change the algorithm
        addresses and a runtime-encoded instruction descriptor
    * - cta_group
      - ``1`` → one CTA, ``M ∈ {64, 128}``; ``2`` → two CTAs split the operand,
-       ``M ∈ {128, 256}`` and half the per-CTA N
+       each with per-CTA ``M ∈ {64, 128}`` and half the B rows. The
+       per-CTA M=64 output uses Layout B
    * - M / N / K extents
      - set the ``(mi, ni, ki)`` unrolled loop counts; K iterations accumulate 
into
        the same tmem accumulator
diff --git a/python/tvm/backend/cuda/lang/alloc_pool.py 
b/python/tvm/backend/cuda/lang/alloc_pool.py
index c7ce8942a6..0dab5341d4 100644
--- a/python/tvm/backend/cuda/lang/alloc_pool.py
+++ b/python/tvm/backend/cuda/lang/alloc_pool.py
@@ -323,7 +323,8 @@ class TMEMPool:
             Explicit ``TileLayout``. Mutually exclusive with ``datapath``.
         datapath : str | None
             Optional tcgen05 datapath letter (``"D"`` for M=128 full datapath,
-            ``"F"`` for M=64 non-``.ws`` scattered). When provided, the 
buffer's
+            ``"F"`` for M=64 non-``.ws`` scattered, or ``"B"`` for per-CTA
+            M=64 ``.cta_group::2`` "2x2"). When provided, the buffer's
             layout is derived from ``tmem_datapath_layout(datapath, *shape)``
             so the row index reflects the *physical* TMEM lane occupation
             (PTX ISA §9.7.16.10.5). The downstream ``.16x*b`` / ``.32x32b``
diff --git 
a/python/tvm/backend/cuda/operator/tile_primitive/copy_async/tcgen05_ldst.py 
b/python/tvm/backend/cuda/operator/tile_primitive/copy_async/tcgen05_ldst.py
index b3842303bb..27144a8e22 100644
--- a/python/tvm/backend/cuda/operator/tile_primitive/copy_async/tcgen05_ldst.py
+++ b/python/tvm/backend/cuda/operator/tile_primitive/copy_async/tcgen05_ldst.py
@@ -88,11 +88,13 @@ def _classify_tmem_datapath(tmem_buf):
     Layout D (M=128, identity row→lane) is the default returned by
     ``_default_tmem_layout``. Layout F (M=64 non-``.ws``, scattered) is the
     explicit opt-in produced by ``tmem_pool.alloc(..., datapath="F")``.
+    Layout B is the per-CTA M=64, ``.cta_group::2`` "2x2" placement produced
+    by ``tmem_pool.alloc(..., datapath="B")``.
     The dispatch uses this to pair each ``.16x*b`` / ``.32x32b`` atom with a
     compatible layout — see ``_check_tmem_layout_for_atom``.
 
-    ``sub_slab`` is always 0 for Layout D. For Layout F it selects the lower
-    (0) or upper (1) 16-lane half of each warp's 32-lane partition.
+    ``sub_slab`` is always 0 for Layout D and B. For Layout F it selects the
+    lower (0) or upper (1) 16-lane half of each warp's 32-lane partition.
     """
     if tmem_buf.layout is None:
         return None
@@ -106,6 +108,15 @@ def _classify_tmem_datapath(tmem_buf):
         except (AssertionError, ValueError):
             return None
     if rows == 64:
+        # Layout B splits N into two N/2 column halves, together spanning all
+        # 128 lanes. Its structure is disjoint from Layout F; try it first.
+        if int(tmem_buf.shape[1]) % 2 == 0:
+            cand = tmem_datapath_layout("B", 64, 
tmem_buf.shape[1]).canonicalize()
+            try:
+                tvm.ir.assert_structural_equal(buf_layout, cand)
+                return ("B", 0)
+            except (AssertionError, ValueError):
+                pass
         # Layout F may occupy either 16-lane half of each warp's 32-lane
         # partition. The layout carries that choice as a +16 TLane offset;
         # thread it through to the PTX row immediate instead of adding an
@@ -142,6 +153,11 @@ def _classify_tmem_datapath(tmem_buf):
 #                                |           |   high slab (row=16) is garbage
 #   F (M=64 scatter)x .32x32b    | no       | F only utilizes 16 of each
 #                                |           |   warp's 32 lanes
+#   B (M=64 2x2)    x bare atom  | no       | B splits N into two N/2
+#                                |           |   lane-halves; use the dedicated
+#                                |           |   logical (64, N) Layout B image
+#                                |           |   handled by
+#                                |           |   _emit_datapath_b_path
 _TMEM_ATOM_COMPAT = {
     ("D", "32x32b", 128): True,
     ("D", "16x*b", 64): True,
@@ -149,6 +165,9 @@ _TMEM_ATOM_COMPAT = {
     ("F", "32x32b", 128): False,
     ("F", "16x*b", 64): True,
     ("F", "16x*b", 128): False,
+    ("B", "32x32b", 128): False,
+    ("B", "16x*b", 64): False,
+    ("B", "16x*b", 128): False,
 }
 
 
@@ -207,6 +226,20 @@ def copy_tmem_local_impl(op_call: TilePrimitiveCall, sctx: 
DispatchContext) -> P
     elem_per_32b = 32 // elem_size
     assert len(local_buf.shape) == len(tmem_buf.shape) == 2
 
+    # Datapath B is identified from the TMEM side before probing ordinary
+    # register atoms. A logical (64, N) Layout B tile is physically a
+    # (128, N/2) .32x32b transfer and therefore needs its dedicated reshape.
+    if _classify_tmem_datapath(tmem_buf) == ("B", 0):
+        return _emit_datapath_b_path(
+            direction=direction,
+            tmem_buf=tmem_buf,
+            local_buf=local_buf,
+            tmem_region=tmem_region,
+            local_region=local_region,
+            elem_per_32b=elem_per_32b,
+            analyzer=analyzer,
+        )
+
     # Try the .16x* (M=64) path first by structural-matching the register-side
     # layout against ``tcgen05_atom_layout(instr_shape, (64, K), dtype)``. The
     # An M=64 TMEM-side Layout F fragment lives in either lanes 0..15 or
@@ -452,6 +485,78 @@ def _emit_16xnb_path(
     return impl
 
 
+def _emit_datapath_b_path(
+    *, direction, tmem_buf, local_buf, tmem_region, local_region, 
elem_per_32b, analyzer
+) -> PrimFunc:
+    """Read or write a Layout B (per-CTA M=64, ``.cta_group::2``) accumulator.
+
+    Layout B splits a logical ``(64, N)`` tile into two ``N/2`` column halves
+    over physical lanes 0..63 and 64..127. It is therefore transferred as one
+    physical ``.32x32b`` ``(128, N/2)`` register file, re-labeled by
+    ``tcgen05_atom_layout("32x32b", (64, N), "float32")``.
+    """
+    if elem_per_32b != 1:
+        raise ValueError(
+            "datapath B readback expects an fp32 fragment (32-bit cells), got "
+            f"dtype={local_buf.dtype!r}"
+        )
+    if int(local_buf.shape[0]) != 64:
+        raise ValueError(
+            "datapath B (.cta_group::2 M=64) fragment must be (64, N); a "
+            f"128-row .32x32b fragment reads the wrong region. Got 
rows={local_buf.shape[0]}. "
+            "Allocate it with T.alloc_tcgen05_ldst_frag('32x32b', (64, N), 
'float32')."
+        )
+
+    n_cols = int(local_buf.shape[1])
+    n_half = n_cols // 2
+    expected_local = tcgen05_atom_layout("32x32b", (64, n_cols), 
local_buf.dtype).canonicalize()
+    try:
+        tvm.ir.assert_structural_equal(local_buf.layout.canonicalize(), 
expected_local)
+    except (AssertionError, ValueError) as err:
+        raise ValueError(
+            "datapath B (.cta_group::2 M=64) requires a matching Layout B "
+            "register fragment. Allocate it with "
+            "T.alloc_tcgen05_ldst_frag('32x32b', (64, N), 'float32'); a "
+            ".16x*b fragment reads the wrong physical lanes and columns. "
+            f"(fragment layout mismatch: {err})"
+        ) from err
+
+    # A partial logical-column slice is not contiguous after the two-way lane
+    # split. Keep this first implementation deliberately strict and transfer
+    # the complete logical tile on both sides.
+    tmem_st, tmem_extent = get_st_extent(tmem_region)
+    local_st, local_extent = get_st_extent(local_region)
+    if not (
+        analyzer.can_prove_equal(tmem_st[0], 0)
+        and analyzer.can_prove_equal(tmem_st[1], 0)
+        and analyzer.can_prove_equal(tmem_extent[0], 64)
+        and analyzer.can_prove_equal(tmem_extent[1], n_cols)
+    ):
+        raise ValueError("datapath B copy must cover the full (64, N) TMEM 
buffer")
+    if not (
+        analyzer.can_prove_equal(local_st[0], 0)
+        and analyzer.can_prove_equal(local_st[1], 0)
+        and analyzer.can_prove_equal(local_extent[0], 64)
+        and analyzer.can_prove_equal(local_extent[1], n_cols)
+    ):
+        raise ValueError("datapath B copy must cover the full (64, N) register 
fragment")
+
+    op = T.ptx.tcgen05.ld if direction == "tmem2local" else T.ptx.tcgen05.st
+
+    # fmt: off
+    @T.prim_func(check_well_formed=False)
+    def impl():
+        local_storage = local_buf.view(n_half, layout=TileLayout(S[n_half]))
+        local_32b = local_storage.view("uint32")
+        op(
+            tmem_buf.allocated_addr[0],
+            *[local_32b[i] for i in range(n_half)],
+            shape="32x32b", num=n_half, row=0, col=0,
+        )
+    # fmt: on
+    return impl
+
+
 # === Variant: copy_async/tmem<->local (priority=10) ===
 #
 # When: one buffer is in tmem (tensor memory, Blackwell SM100+) and the other
diff --git 
a/python/tvm/backend/cuda/operator/tile_primitive/gemm_async/tcgen05.py 
b/python/tvm/backend/cuda/operator/tile_primitive/gemm_async/tcgen05.py
index 9d0602234d..637ca13787 100644
--- a/python/tvm/backend/cuda/operator/tile_primitive/gemm_async/tcgen05.py
+++ b/python/tvm/backend/cuda/operator/tile_primitive/gemm_async/tcgen05.py
@@ -776,8 +776,10 @@ def gemm_async_tcgen05_impl(op_call: TilePrimitiveCall, 
sctx: DispatchContext) -
     # the full row range, the slice layout structurally matches Layout F over
     # (M=64, N) — assert against that base instead of the Layout D identity.
     if is_2x2:
-        N_half = N // 2
-        base = TileLayout(S[(M, 2, N_half) : (1 @ TLane, 64 @ TLane, 1 @ 
TCol)])
+        # Layout B (per-CTA M=64, .cta_group::2 "2x2"). Single-source the
+        # write layout with allocation and readback so the two sides cannot
+        # drift.
+        base = tmem_datapath_layout("B", M, N)
     elif (
         M == 64
         and int(C_buffer.shape[0]) == 64
diff --git a/python/tvm/tirx/layout.py b/python/tvm/tirx/layout.py
index 3557e35367..3dcef98bc4 100644
--- a/python/tvm/tirx/layout.py
+++ b/python/tvm/tirx/layout.py
@@ -585,21 +585,27 @@ __all__ += ["tcgen05_atom_layout", 
"tmem_datapath_layout", "wg_local_layout"]
 # scatter directly.
 #
 # We surface this via the factory below. Callers pass the datapath letter
-# (``"D"`` / ``"F"``) and the logical ``(rows, cols)``; the factory returns
-# the appropriate TileLayout. ``tmem_pool.alloc(..., datapath="F")`` plumbs
-# this into the buffer's layout so the dispatch can structurally verify
-# atom ↔ datapath compatibility instead of silently accepting mismatches.
+# (``"D"`` / ``"F"`` / ``"B"``) and the logical ``(rows, cols)``; the
+# factory returns the appropriate TileLayout. ``tmem_pool.alloc(...,
+# datapath="F")`` plumbs this into the buffer's layout so the dispatch can
+# structurally verify atom ↔ datapath compatibility instead of silently
+# accepting mismatches.
 #
 # Supported today:
 #   - ``"D"``: M=128, ``.cta_group::1``, full datapath. Identity row→lane.
 #   - ``"F"``: M=64, non-``.ws``, half datapath (4x1 lane utilization).
 #     Logical row r → physical lane
 #     (r // 16) * 32 + sub_slab * 16 + (r % 16).
+#   - ``"B"``: per-CTA M=64, ``.cta_group::2``, Dense A ("2x2" datapath).
+#     PTX names the CTA-pair shape M=128; each CTA owns a logical ``(64, N)``
+#     accumulator. Its N columns split into two N/2 halves across physical
+#     lanes 0..63 and 64..127:
+#       (r, c) → (TLane=r + 64 * (c // (N/2)), TCol=c % (N/2)).
 #
-# Layouts A / B / C / E / G are reserved for future expansion.
+# Layouts A / C / E / G are reserved for future expansion.
 
 
-_TMEM_DATAPATH_ROWS = {"D": 128, "F": 64}
+_TMEM_DATAPATH_ROWS = {"D": 128, "F": 64, "B": 64}
 
 
 def tmem_datapath_layout(datapath: str, rows: int, cols: int, sub_slab: int = 
0) -> "TileLayout":
@@ -614,20 +620,22 @@ def tmem_datapath_layout(datapath: str, rows: int, cols: 
int, sub_slab: int = 0)
     Parameters
     ----------
     datapath : str
-        One of ``"D"`` (M=128, ``.cta_group::1``, full datapath) or
-        ``"F"`` (M=64, non-``.ws``, half datapath). Other layouts are not
-        yet supported by this factory.
+        One of ``"D"`` (M=128, ``.cta_group::1``, full datapath), ``"F"``
+        (M=64, non-``.ws``, half datapath), or ``"B"`` (per-CTA M=64,
+        ``.cta_group::2``, "2x2" datapath). Other layouts are not yet
+        supported by this factory.
     rows : int
         Logical row count of the TMEM buffer. Must match the datapath's M
-        dimension: 128 for D, 64 for F.
+        dimension: 128 for D, 64 for F and B.
     cols : int
-        Logical column count.
+        Logical column count. Datapath B requires an even count because its
+        columns split into two equal lane halves.
     sub_slab : int
         For Layout F, select the lower (``0``) or upper (``1``) 16-lane
         half of each warp's 32-lane TMEM partition. The upper half is useful
         as a 64-row read/write view of the high half-slab of a Layout D
-        accumulator. Layout D already spans both halves and therefore only
-        accepts ``0``.
+        accumulator. Layouts D and B already span both halves and therefore
+        only accept ``0``.
 
     Returns
     -------
@@ -657,6 +665,22 @@ def tmem_datapath_layout(datapath: str, rows: int, cols: 
int, sub_slab: int = 0)
                 "sub-slabs; sub_slab must be 0"
             )
         return TileLayout(S[(rows, cols) : (1 @ tlane, 1 @ tcol)])
+    if datapath == "B":
+        # Layout B: a per-CTA (64, N) accumulator uses the same physical
+        # footprint as a Layout D (128, N/2) accumulator. The high column bit
+        # selects one of the two 64-lane halves, while the low bits select 
TCol.
+        # B always spans all 128 lanes, so the F-only sub_slab selector does 
not
+        # apply.
+        if sub_slab != 0:
+            raise ValueError(
+                "tmem_datapath_layout: datapath='B' spans all 128 lanes; 
sub_slab must be 0"
+            )
+        if cols % 2 != 0:
+            raise ValueError(
+                f"tmem_datapath_layout: datapath='B' expects even cols (N), 
got {cols}"
+            )
+        n_half = cols // 2
+        return TileLayout(S[(rows, 2, n_half) : (1 @ tlane, 64 @ tlane, 1 @ 
tcol)])
     # Layout F: M=64 scattered. Logical row r = wid * 16 + intra (wid ∈ [0,4),
     # intra ∈ [0,16)) → physical lane wid * 32 + sub_slab * 16 + intra.
     # ``TileLayout`` decomposes a scalar row index via ``SplitCoord``
@@ -697,14 +721,14 @@ _TCGEN05_ATOM_REPS = {
 _TCGEN05_COL_FACTOR_FP32 = {"32x32b": 1, "16x64b": 2, "16x128b": 4, "16x256b": 
8}
 
 # Allowed fragment row counts per warpgroup for each instr_shape. ``.32x32b``
-# is fixed at M=128; ``.16x*b`` natively covers M=64 (one 16-row slab per
-# warp, using lanes 0..15 of each warp's 32-lane TMEM partition) and can be
-# extended to M=128 by issuing the atom twice with row offsets 0 and 16
-# (covering lanes 0..15 + 16..31, i.e. the warp's full slab). The M=128
-# variant doubles per-thread registers and treats the extra slab as the
-# highest m-bit.
+# normally covers M=128, but a 64-row shape denotes the Layout B readback 
image:
+# the logical (64, N) tile is physically a (128, N/2) ``.32x32b`` register
+# file. ``.16x*b`` natively covers M=64 (one 16-row slab per warp, using lanes
+# 0..15 of each warp's 32-lane TMEM partition) and can be extended to M=128 by
+# issuing the atom twice with row offsets 0 and 16. The M=128 variant doubles
+# per-thread registers and treats the extra slab as the highest m-bit.
 _TCGEN05_FRAG_ROWS = {
-    "32x32b": (128,),
+    "32x32b": (64, 128),
     "16x64b": (64, 128),
     "16x128b": (64, 128),
     "16x256b": (64, 128),
@@ -712,15 +736,17 @@ _TCGEN05_FRAG_ROWS = {
 
 
 def tcgen05_atom_layout(instr_shape: str, tensor_shape: tuple[int, int], 
dtype) -> "TileLayout":
-    """Register-side ``TileLayout`` for ``tcgen05.ld``/``tcgen05.st`` 
``.16x*`` atoms.
+    """Register-side ``TileLayout`` for ``tcgen05.ld``/``tcgen05.st`` atoms.
 
     Describes the per-warpgroup register tile that ``Tx.copy_async`` produces
     when reading a TMEM fragment via ``tcgen05.{ld,st}.<instr_shape>.xN``.
     ``rep`` (the ``.xN`` qualifier) is inferred from ``tensor_shape``.
 
-    Fragment row count is determined by ``instr_shape``: ``.32x32b`` covers an
-    M=128 fragment (128 rows per warpgroup), and ``.16x{64,128,256}b`` covers
-    an M=64 fragment (64 rows per warpgroup).
+    Fragment row count is determined by ``instr_shape``: ``.32x32b`` normally
+    covers an M=128 fragment (128 rows per warpgroup), and
+    ``.16x{64,128,256}b`` covers an M=64 fragment (64 rows per warpgroup).
+    ``("32x32b", (64, N))`` is the fp32 Layout B readback image for a
+    ``.cta_group::2`` M=64 accumulator.
 
     TMEM is kept **dense** for 16-bit dtypes: two 16-bit elements per 32-bit
     TMEM cell (matching the existing ``.32x32b`` convention). The PTX op is
@@ -736,8 +762,9 @@ def tcgen05_atom_layout(instr_shape: str, tensor_shape: 
tuple[int, int], dtype)
     tensor_shape : tuple[int, int]
         The logical fragment shape in **element units**. Must be
         ``(frag_rows, K)`` where ``frag_rows`` is ``128`` for ``.32x32b`` and
-        ``64`` for the other shapes, and ``K`` is divisible by the per-warp
-        column factor for the chosen instr_shape and dtype::
+        ``64`` for the other shapes. The fp32 Layout B image is the exception:
+        it uses ``("32x32b", (64, N))``. The column extent is divisible by the
+        per-warp column factor for the chosen instr_shape and dtype::
 
             K must be a power-of-two multiple of (factor_fp32 * elem_per_32b)
 
@@ -751,9 +778,9 @@ def tcgen05_atom_layout(instr_shape: str, tensor_shape: 
tuple[int, int], dtype)
     Returns
     -------
     TileLayout
-        A ``(64, K)``-shaped tile layout. The factory builds it as a sequence
-        of fine-grained iters describing the per-(lane, register) destination
-        position; ``.group([(64, K)])[0]`` flattens to two iters.
+        A ``tensor_shape``-shaped tile layout. The factory builds it as a
+        sequence of fine-grained iters describing the per-(lane, register)
+        destination position.
 
     Examples
     --------
@@ -783,6 +810,35 @@ def tcgen05_atom_layout(instr_shape: str, tensor_shape: 
tuple[int, int], dtype)
             f"tcgen05_atom_layout {instr_shape!r} expects rows ∈ 
{allowed_rows}, got {rows}"
         )
 
+    if instr_shape == "32x32b" and rows == 64:
+        # Layout B's logical (64, N) tile physically occupies all 128 lanes and
+        # N/2 tcols. A .32x32b thread therefore owns its physical lane and N/2
+        # fp32 registers. Re-label that (128, N/2) register file as (64, N).
+        if bits != 32:
+            raise ValueError(
+                "tcgen05_atom_layout: 32x32b with 64 rows is the datapath B "
+                f"readback image and is fp32-only, got {dtype} ({bits} bits)"
+            )
+        if cols % 2 != 0:
+            raise ValueError(
+                f"tcgen05_atom_layout: 32x32b (64, N) datapath B image expects 
even N, got {cols}"
+            )
+        n_half = cols // 2
+        if n_half not in _TCGEN05_ATOM_REPS["32x32b"]:
+            raise ValueError(
+                f"tcgen05_atom_layout: 32x32b (64, N) datapath B image needs 
N/2={n_half} "
+                f"(from N={cols}) in the PTX Table 49 set 
{_TCGEN05_ATOM_REPS['32x32b']}"
+            )
+        return TileLayout.from_iters(
+            [
+                Iter(64, 1, Axis.tid_in_wg),
+                Iter(2, 64, Axis.tid_in_wg),
+                Iter(n_half, 1, "m"),
+            ],
+            [],
+            {},
+        )
+
     elem_per_32b = 32 // bits
     col_factor_elem = _TCGEN05_COL_FACTOR_FP32[instr_shape] * elem_per_32b
     if cols % col_factor_elem != 0:
diff --git a/python/tvm/tirx/script/builder/ir.py 
b/python/tvm/tirx/script/builder/ir.py
index 43a383a1d2..4041ef8386 100644
--- a/python/tvm/tirx/script/builder/ir.py
+++ b/python/tvm/tirx/script/builder/ir.py
@@ -1807,7 +1807,7 @@ def alloc_tcgen05_ldst_frag(instr_shape, tensor_shape, 
dtype):
 
     Sizes the per-thread storage, allocates ``local`` scope memory, and returns
     a 2-D view of shape ``tensor_shape`` with a matching 
``tcgen05_atom_layout``.
-    Pass the result to ``Tx.copy_async`` (with a ``(128, W)``-shaped TMEM
+    Pass the result to ``Tx.wg.copy_async`` (with a matching TMEM
     buffer) to trigger the corresponding dispatch path.
 
     Parameters
@@ -1819,7 +1819,8 @@ def alloc_tcgen05_ldst_frag(instr_shape, tensor_shape, 
dtype):
         per-shape per-lane register decomposition).
     tensor_shape : tuple[int, int]
         Logical fragment shape ``(frag_rows, K)`` in element units. 
``frag_rows``
-        is ``128`` for ``.32x32b`` and ``64`` for the ``.16x*b`` shapes.
+        is ``128`` for ``.32x32b`` and ``64`` for the ``.16x*b`` shapes. The
+        fp32 Layout B readback image also uses ``("32x32b", (64, N))``.
     dtype : str
         ``"float32"``, ``"float16"``, or ``"bfloat16"``.
 
@@ -1833,11 +1834,16 @@ def alloc_tcgen05_ldst_frag(instr_shape, tensor_shape, 
dtype):
     --------
     M=128 readback (existing dispatch):
         ``frag = T.alloc_tcgen05_ldst_frag("32x32b", (128, 64), "float32")``
-        ``Tx.copy_async(frag[:, :], tmem[:, 0:64])``
+        ``Tx.wg.copy_async(frag[:, :], tmem[:, 0:64])``
 
     M=64 readback (.16x64b dispatch):
         ``frag = T.alloc_tcgen05_ldst_frag("16x64b", (64, 64), "float32")``
-        ``Tx.copy_async(frag[:, :], tmem[0:64, 0:64])``
+        ``Tx.wg.copy_async(frag[:, :], tmem[0:64, 0:64])``
+
+    Datapath B readback (cta_group=2, per-CTA M=64):
+        ``C = tmem_pool.alloc((64, 128), "float32", datapath="B")``
+        ``frag = T.alloc_tcgen05_ldst_frag("32x32b", (64, 128), "float32")``
+        ``Tx.wg.copy_async(frag[:, :], C[:, :])``
     """
     from tvm.tirx.layout import tcgen05_atom_layout  # local import to avoid 
cycle
 
diff --git 
a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tmem_16xnb.py 
b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tmem_16xnb.py
index 31c159dc94..ac35f4c628 100644
--- 
a/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tmem_16xnb.py
+++ 
b/tests/python/tirx/operator/tile_primitive/cuda/copy_async/test_tmem_16xnb.py
@@ -452,6 +452,47 @@ def test_tmem_datapath_layout_D_row_to_lane_mapping():
         )
 
 
+def test_tmem_datapath_layout_B_col_split_mapping():
+    """Layout B splits N/2 columns across each 64-lane half."""
+    n_cols = 128
+    n_half = n_cols // 2
+    layout = tmem_datapath_layout("B", 64, n_cols)
+
+    for row in [0, 1, 31, 63]:
+        for col in [0, 1, n_half - 1, n_half, n_half + 1, n_cols - 1]:
+            axis_values = layout.apply(row, col, shape=[64, n_cols])
+            assert int(axis_values["TLane"]) == row + 64 * (col // n_half)
+            assert int(axis_values["TCol"]) == col % n_half
+
+
+def test_tcgen05_atom_layout_32x32b_datapath_B_mapping():
+    """The Layout B register image mirrors TLane/TCol as tid_in_wg/m."""
+    n_cols = 128
+    n_half = n_cols // 2
+    layout = tcgen05_atom_layout("32x32b", (64, n_cols), "float32")
+
+    for row in [0, 1, 31, 63]:
+        for col in [0, 1, n_half - 1, n_half, n_half + 1, n_cols - 1]:
+            axis_values = layout.apply(row, col, shape=[64, n_cols])
+            assert int(axis_values["tid_in_wg"]) == row + 64 * (col // n_half)
+            assert int(axis_values["m"]) == col % n_half
+
+
+def test_datapath_B_layout_factories_reject_invalid_inputs():
+    with pytest.raises(ValueError, match="expects rows=64"):
+        tmem_datapath_layout("B", 128, 32)
+    with pytest.raises(ValueError, match="expects even cols"):
+        tmem_datapath_layout("B", 64, 31)
+    with pytest.raises(ValueError, match="sub_slab must be 0"):
+        tmem_datapath_layout("B", 64, 32, sub_slab=1)
+    with pytest.raises(ValueError, match="fp32-only"):
+        tcgen05_atom_layout("32x32b", (64, 32), "float16")
+    with pytest.raises(ValueError, match="expects even N"):
+        tcgen05_atom_layout("32x32b", (64, 31), "float32")
+    with pytest.raises(ValueError, match="PTX Table 49"):
+        tcgen05_atom_layout("32x32b", (64, 6), "float32")
+
+
 @pytest.mark.gpu
 @pytest.mark.skipif(not env.has_cuda_compute(10), reason="need cuda compute >= 
10.0")
 @pytest.mark.parametrize("shape,rep", [("16x256b", 4), ("16x128b", 4), 
("16x64b", 8)])
@@ -614,6 +655,169 @@ def test_layout_F_rejects_incompatible_atoms(atom_kind, 
frag_rows):
             tvm.compile(mod, target=target, tir_pipeline="tirx")
 
 
+def test_layout_B_rejects_16xnb_fragment():
+    """Layout B must not silently take the ordinary M=64 .16x*b path."""
+    n_cols = 128
+    tmem_layout = tmem_datapath_layout("B", 64, n_cols)
+    wrong_layout = tcgen05_atom_layout("16x256b", (64, n_cols), "float32")
+
+    @T.prim_func
+    def kernel() -> None:
+        T.device_entry()
+        T.cta_id([1])
+        T.warpgroup_id([1])
+        T.warp_id_in_wg([4])
+        T.lane_id([32])
+        tmem_addr = T.alloc_shared([1], "uint32")
+        tmem = T.decl_buffer(
+            (64, n_cols),
+            "float32",
+            scope="tmem",
+            allocated_addr=tmem_addr[0],
+            layout=tmem_layout,
+        )
+        frag = T.alloc_local((n_cols // 2,), "float32")
+        frag_view = frag.view(64, n_cols, layout=wrong_layout)
+        Tx.wg.copy_async(frag_view[:, :], tmem[:, :])
+
+    target = tvm.target.Target("cuda")
+    with target:
+        with pytest.raises((ValueError, RuntimeError), match="datapath B"):
+            tvm.compile(tvm.IRModule({"main": kernel}), target=target, 
tir_pipeline="tirx")
+
+
+def test_layout_B_rejects_partial_column_copy():
+    """A logical B column slice is not one contiguous physical tcol 
interval."""
+    n_cols = 64
+
+    @T.prim_func
+    def kernel() -> None:
+        T.device_entry()
+        T.cta_id([1])
+        T.warpgroup_id([1])
+        T.warp_id_in_wg([4])
+        T.lane_id([32])
+        tmem_addr = T.alloc_shared([1], "uint32")
+        tmem = T.decl_buffer(
+            (64, n_cols),
+            "float32",
+            scope="tmem",
+            allocated_addr=tmem_addr[0],
+            layout=tmem_datapath_layout("B", 64, n_cols),
+        )
+        frag = T.alloc_tcgen05_ldst_frag("32x32b", (64, n_cols), "float32")
+        Tx.wg.copy_async(frag[:, : n_cols // 2], tmem[:, : n_cols // 2])
+
+    target = tvm.target.Target("cuda")
+    with target:
+        with pytest.raises((ValueError, RuntimeError), match=r"full \(64, 
N\)"):
+            tvm.compile(tvm.IRModule({"main": kernel}), target=target, 
tir_pipeline="tirx")
+
+
[email protected]("direction", ["ld", "st"])
+def test_datapath_B_codegen(direction):
+    """Both directions emit one physical .32x32b.x32 instruction."""
+    n_cols = 64
+
+    @T.prim_func
+    def kernel() -> None:
+        T.device_entry()
+        T.cta_id([1])
+        T.warpgroup_id([1])
+        T.warp_id_in_wg([4])
+        T.lane_id([32])
+        tmem_addr = T.alloc_shared([1], "uint32")
+        tmem = T.decl_buffer(
+            (64, n_cols),
+            "float32",
+            scope="tmem",
+            allocated_addr=tmem_addr[0] + 32,
+            layout=tmem_datapath_layout("B", 64, n_cols),
+        )
+        frag = T.alloc_tcgen05_ldst_frag("32x32b", (64, n_cols), "float32")
+        if direction == "ld":
+            Tx.wg.copy_async(frag[:, :], tmem[:, :])
+        else:
+            Tx.wg.copy_async(tmem[:, :], frag[:, :])
+
+    target = tvm.target.Target("cuda")
+    with target:
+        mod = tvm.compile(tvm.IRModule({"main": kernel}), target=target, 
tir_pipeline="tirx")
+    source = mod.mod.imports[0].inspect_source()
+    assert f"tcgen05.{direction}" in source
+    assert "32x32b.x32" in source
+
+
[email protected]
[email protected](not env.has_cuda_compute(10), reason="need cuda compute >= 
10.0")
[email protected]("n_cols", [32, 64, 128, 256])
[email protected]("col_offset", [0, 32])
+def test_datapath_B_ld_st_roundtrip(n_cols, col_offset):
+    """Layout B store/load preserves every register, including a nonzero 
base."""
+    n_half = n_cols // 2
+    tmem_cols = _next_pow2(max(32, col_offset + n_half))
+
+    @T.prim_func
+    def kernel(A_ptr: T.handle, B_ptr: T.handle) -> None:
+        A = T.match_buffer(A_ptr, (128, n_half), "float32")
+        B = T.match_buffer(B_ptr, (128, n_half), "float32")
+        T.device_entry()
+        warp_id = T.warp_id([4])
+        T.cta_id([1])
+        wg_id = T.warpgroup_id([1])
+        T.warp_id_in_wg([4])
+        T.lane_id([32])
+        tid = T.thread_id_in_wg([128])
+        tmem_addr = T.alloc_shared([1], "uint32")
+
+        if wg_id == 0:
+            if warp_id == 0:
+                T.ptx.tcgen05.alloc(T.address_of(tmem_addr), n_cols=tmem_cols, 
cta_group=1)
+            T.tvm_storage_sync("shared")
+            tmem = T.decl_buffer(
+                (64, n_cols),
+                "float32",
+                scope="tmem",
+                allocated_addr=tmem_addr[0] + col_offset,
+                layout=tmem_datapath_layout("B", 64, n_cols),
+            )
+
+            frag_in = T.alloc_tcgen05_ldst_frag("32x32b", (64, n_cols), 
"float32")
+            frag_in_local = frag_in.local()
+            for i in range(n_half):
+                frag_in_local[i] = A[tid, i]
+            T.cuda.cta_sync()
+            Tx.wg.copy_async(tmem[:, :], frag_in[:, :])
+            T.ptx.tcgen05.wait.st()
+            T.cuda.cta_sync()
+
+            frag_out = T.alloc_tcgen05_ldst_frag("32x32b", (64, n_cols), 
"float32")
+            Tx.wg.copy_async(frag_out[:, :], tmem[:, :])
+            T.ptx.tcgen05.wait.ld()
+            T.cuda.cta_sync()
+            frag_out_local = frag_out.local()
+            for i in range(n_half):
+                B[tid, i] = frag_out_local[i]
+
+            if warp_id == 0:
+                T.ptx.tcgen05.relinquish_alloc_permit(cta_group=1)
+                T.ptx.tcgen05.dealloc(tmem_addr[0], n_cols=tmem_cols, 
cta_group=1)
+
+    target = tvm.target.Target("cuda")
+    with target:
+        mod = tvm.compile(tvm.IRModule({"main": kernel}), target=target, 
tir_pipeline="tirx")
+        source_np = tvm.testing.generate_random_array("float32", (128, n_half))
+
+        def run_and_check():
+            dev = tvm.cuda(0)
+            source = tvm.runtime.tensor(source_np, dev)
+            result = tvm.runtime.tensor(np.zeros((128, n_half), 
dtype="float32"), dev)
+            mod(source, result)
+            np.testing.assert_array_equal(result.numpy(), source_np)
+
+        tvm.testing.run_with_gpu_lock(run_and_check)
+
+
 def _run_load_test(shape: str, rep: int, dtype: str):
     """Stage A into TMEM via .32x32b, then read it back as the fragment via
     .<shape>.x<rep> (through ``T.alloc_tcgen05_ldst_frag``), and compare each
diff --git 
a/tests/python/tirx/operator/tile_primitive/cuda/gemm_async/test_gemm_async.py 
b/tests/python/tirx/operator/tile_primitive/cuda/gemm_async/test_gemm_async.py
index d4b8917436..716c0c02e3 100644
--- 
a/tests/python/tirx/operator/tile_primitive/cuda/gemm_async/test_gemm_async.py
+++ 
b/tests/python/tirx/operator/tile_primitive/cuda/gemm_async/test_gemm_async.py
@@ -39,7 +39,14 @@ from tvm.tirx.cuda.operator.tile_primitive.tma_utils import (
     mma_atom_shape,
     mma_shared_layout,
 )
-from tvm.tirx.layout import S, TCol, TileLayout, TLane, tcgen05_atom_layout
+from tvm.tirx.layout import (
+    S,
+    TCol,
+    TileLayout,
+    TLane,
+    tcgen05_atom_layout,
+    tmem_datapath_layout,
+)
 from tvm.tirx.layout import tid_in_wg as axis_tid_in_wg
 
 # ---------------------------------------------------------------------------
@@ -483,7 +490,7 @@ def test_gemm_tcgen05_cta_group_2(task):
         tma_mbar = T.alloc_shared([1], "uint64")
         mma_mbar = T.alloc_shared([1], "uint64")
 
-        ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64")))] = 
T.reinterpret("handle", T.ptx.map_shared_rank(tma_mbar.ptr_to([0]), 0))  # 
noqa: E501
+        ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64"), 
"shared"))] = T.reinterpret(PointerType(PrimType("uint64"), "shared"), 
T.ptx.map_shared_rank(tma_mbar.ptr_to([0]), 0))  # noqa: E501
         tma_mbar_cta_0 = T.decl_buffer([1], "uint64", data=ptr, scope="shared")
 
         if tid_in_wg == 0:
@@ -615,7 +622,7 @@ def test_gemm_tcgen05_cta_group_2_layout_b():
         tma_mbar = T.alloc_shared([1], "uint64")
         mma_mbar = T.alloc_shared([1], "uint64")
 
-        ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64")))] = 
T.reinterpret("handle", T.ptx.map_shared_rank(tma_mbar.ptr_to([0]), 0))  # 
noqa: E501
+        ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64"), 
"shared"))] = T.reinterpret(PointerType(PrimType("uint64"), "shared"), 
T.ptx.map_shared_rank(tma_mbar.ptr_to([0]), 0))  # noqa: E501
         tma_mbar_cta_0 = T.decl_buffer([1], "uint64", data=ptr, scope="shared")
 
         if tid_in_wg == 0:
@@ -694,6 +701,134 @@ def test_gemm_tcgen05_cta_group_2_layout_b():
         tvm.testing.run_with_gpu_lock(run_and_check)
 
 
[email protected]
[email protected](not env.has_cuda_compute(10), reason="need cuda compute >= 
10.0")
+def test_gemm_tcgen05_cta_group_2_datapath_b_readback():
+    """A cta_group=2 GEMM writes and reads a first-class datapath B buffer."""
+    m_per_cta = 64
+    n_logical = 128
+    n_per_cta = n_logical // 2
+    k = 64
+    input_dtype = "float32"
+    a_dtype = "float16"
+    b_dtype = "float16"
+    c_dtype = "float32"
+
+    a_shape = (m_per_cta, k)
+    b_shape = (n_per_cta, k)
+    c_shape = (m_per_cta * 2, n_logical)
+    a_layout = mma_shared_layout(a_dtype, 3, a_shape)
+    b_layout = mma_shared_layout(b_dtype, 3, b_shape)
+
+    # fmt: off
+    @T.prim_func
+    def gemm_async(A_ptr: T.handle, B_ptr: T.handle, C_ptr: T.handle) -> None:
+        A = T.match_buffer(A_ptr, (m_per_cta * 2, k), input_dtype)
+        B = T.match_buffer(B_ptr, (n_logical, k), input_dtype)
+        C = T.match_buffer(C_ptr, c_shape, c_dtype)
+
+        T.device_entry()
+        warp_id = T.warp_id([4])
+        cbx, cby = T.cta_id_in_cluster([2, 1])
+        T.cta_id([2])
+        wg_id = T.warpgroup_id([1])
+        tid = T.thread_id_in_wg([128])
+
+        A_smem = T.alloc_buffer(a_shape, a_dtype, scope="shared", 
layout=a_layout)
+        B_smem = T.alloc_buffer(b_shape, b_dtype, scope="shared", 
layout=b_layout)
+        tmem_addr = T.alloc_shared([1], "uint32")
+        mma_mbar = T.alloc_shared([1], "uint64")
+
+        if tid == 0:
+            T.ptx.mbarrier.init(mma_mbar.ptr_to([0]), 1)
+        if warp_id == 0:
+            T.ptx.tcgen05.alloc(T.address_of(tmem_addr), n_cols=n_per_cta, 
cta_group=2)
+        T.ptx.fence.mbarrier_init()
+        T.cuda.cta_sync()
+
+        tmem = T.decl_buffer(
+            (m_per_cta, n_logical),
+            c_dtype,
+            scope="tmem",
+            allocated_addr=tmem_addr[0],
+            layout=tmem_datapath_layout("B", m_per_cta, n_logical),
+        )
+
+        # Use ordinary shared-memory stores here so this test is independent
+        # of the TMA/remote-mbarrier path exercised by the older Layout B test.
+        for i in range(m_per_cta * k // 128):
+            A_smem[(tid + i * 128) // k, (tid + i * 128) % k] = 
T.Cast(a_dtype, A[
+                cbx * m_per_cta + (tid + i * 128) // k, (tid + i * 128) % k
+            ])
+        for i in range(n_per_cta * k // 128):
+            B_smem[(tid + i * 128) // k, (tid + i * 128) % k] = 
T.Cast(b_dtype, B[
+                cbx * n_per_cta + (tid + i * 128) // k, (tid + i * 128) % k
+            ])
+        T.cuda.cta_sync()
+        T.ptx.fence.proxy_async("shared::cta")
+        T.cuda.cluster_sync()
+
+        if cbx == 0:
+            T.ptx.tcgen05.fence.after_thread_sync()
+            T.cuda.cta_sync()
+            if tid == 0:
+                Tx.gemm_async(
+                    tmem[:, :],
+                    A_smem[:, :],
+                    B_smem[:, :],
+                    dispatch="tcgen05",
+                    cta_group=2,
+                )
+                T.ptx.tcgen05.commit(mma_mbar.ptr_to([0]), cta_group=2, 
cta_mask=3)
+
+        T.ptx.mbarrier.try_wait(mma_mbar.ptr_to([0]), 0)
+        T.ptx.tcgen05.fence.after_thread_sync()
+        T.cuda.cta_sync()
+
+        frag = T.alloc_tcgen05_ldst_frag(
+            "32x32b", (m_per_cta, n_logical), c_dtype
+        )
+        if wg_id == 0:
+            Tx.wg.copy_async(frag[:, :], tmem[:, :])
+            T.ptx.tcgen05.wait.ld()
+        T.cuda.cta_sync()
+
+        frag_local = frag.local()
+        for i in range(n_per_cta):
+            C[
+                cbx * m_per_cta + tid % m_per_cta,
+                (tid // m_per_cta) * n_per_cta + i,
+            ] = frag_local[i]
+        T.cuda.cta_sync()
+
+        if warp_id == 0:
+            T.ptx.tcgen05.relinquish_alloc_permit(cta_group=2)
+            T.ptx.tcgen05.dealloc(tmem_addr[0], n_cols=n_per_cta, cta_group=2)
+        # fmt: on
+
+    target = tvm.target.Target("cuda")
+    with target:
+        mod = tvm.compile(
+            tvm.IRModule({"main": gemm_async}), target=target, 
tir_pipeline="tirx"
+        )
+
+    np.random.seed(0)
+    a_np = np.random.randn(m_per_cta * 2, k).astype(input_dtype)
+    b_np = np.random.randn(n_logical, k).astype(input_dtype)
+    c_np = np.zeros(c_shape, dtype=c_dtype)
+    c_ref = a_np.astype(a_dtype).astype(np.float32) @ 
b_np.astype(b_dtype).astype(np.float32).T
+
+    def run_and_check():
+        dev = tvm.cuda(0)
+        a_tvm = tvm.runtime.tensor(a_np, dev)
+        b_tvm = tvm.runtime.tensor(b_np, dev)
+        c_tvm = tvm.runtime.tensor(c_np, dev)
+        mod["main"](a_tvm, b_tvm, c_tvm)
+        np.testing.assert_allclose(c_tvm.numpy(), c_ref, atol=1e-3, rtol=1e-3)
+
+    tvm.testing.run_with_gpu_lock(run_and_check)
+
+
 @pytest.mark.gpu
 @pytest.mark.skipif(not env.has_cuda_compute(10), reason="need cuda compute >= 
10.0")
 @pytest.mark.skipif(ml_dtypes is None, reason="Requires ml_dtypes")
@@ -992,7 +1127,7 @@ def test_gemm_block_scaled_fp8_cta_group_2(task):
         descSFA = T.alloc_buffer((1,), "uint64", scope="local")
         descSFB = T.alloc_buffer((1,), "uint64", scope="local")
 
-        ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64")))] = 
T.reinterpret("handle", T.ptx.map_shared_rank(tma_mbar.ptr_to([0]), 0))  # 
noqa: E501
+        ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64"), 
"shared"))] = T.reinterpret(PointerType(PrimType("uint64"), "shared"), 
T.ptx.map_shared_rank(tma_mbar.ptr_to([0]), 0))  # noqa: E501
         tma_mbar_cta_0 = T.decl_buffer([1], "uint64", data=ptr, scope="shared")
 
         if tid_in_wg == 0:
@@ -1374,7 +1509,7 @@ def test_gemm_block_scaled_nvfp4_cta_group_2():
         descSFA = T.alloc_buffer((1,), "uint64", scope="local")
         descSFB = T.alloc_buffer((1,), "uint64", scope="local")
 
-        ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64")))] = 
T.reinterpret("handle", T.ptx.map_shared_rank(tma_mbar.ptr_to([0]), 0))  # 
noqa: E501
+        ptr: T.let[T.Var(name="ptr", ty=PointerType(PrimType("uint64"), 
"shared"))] = T.reinterpret(PointerType(PrimType("uint64"), "shared"), 
T.ptx.map_shared_rank(tma_mbar.ptr_to([0]), 0))  # noqa: E501
         tma_mbar_cta_0 = T.decl_buffer([1], "uint64", data=ptr, scope="shared")
 
         if tid_in_wg == 0:

Reply via email to