@@ -66,14 +66,16 @@ class PTXAsyncCopyInjector : public StmtMutator {
6666 }
6767
6868 Stmt VisitStmt_ (const ForNode *op) final {
69- // Track nested vectorized loop extents so we can rewrite element-wise
70- // copies (e.g. float16 stores) into `tir.ptx_cp_async` with element bytes,
71- // relying on the later `tl.VectorizeLoop` pass to widen:
72- // for v in T.vectorized(k): ptx_cp_async(dst, src, elem_bytes)
73- // => ptx_cp_async(dst_base, src_base, elem_bytes * k)
69+ // Track nested vectorized loop extents so cp.async injection can emit the
70+ // final packed transfer width directly:
71+ // for v in T.vectorized(k): store(load(...))
72+ // => ptx_cp_async(dst_base, src_base, total_transfer_bytes)
7473 //
75- // This mirrors the logic in `CPAsyncStoreRewriter` used by `T.copy`
76- // lowering, and avoids duplicating vectorize-loop collapse here.
74+ // cp.async only supports byte widths {4, 8, 16}. Instead of emitting
75+ // element-sized transfers here and relying on a later vectorization pass
76+ // to widen them, TryInjectPTX packs the active vectorized loop directly
77+ // whenever the overall transfer width is legal, then collapses the now
78+ // redundant vectorized loop.
7779 int previous_vectorized_lanes = current_vectorized_lanes_;
7880 bool pushed_vectorized_loop = false ;
7981 if (op->kind == ForKind::kVectorized ) {
@@ -90,6 +92,13 @@ class PTXAsyncCopyInjector : public StmtMutator {
9092 }
9193 }
9294 Stmt stmt = StmtMutator::VisitStmt_ (op);
95+ if (pushed_vectorized_loop) {
96+ if (const auto *loop = stmt.as <ForNode>()) {
97+ if (CanCollapseVectorizedCPAsyncLoop (loop->body , loop->loop_var )) {
98+ stmt = loop->body ;
99+ }
100+ }
101+ }
93102 if (pushed_vectorized_loop) {
94103 active_vectorized_loops_.pop_back ();
95104 }
@@ -123,11 +132,29 @@ class PTXAsyncCopyInjector : public StmtMutator {
123132 index_info->dst_index )) {
124133 return Optional<Stmt>();
125134 }
135+ if (index_info->collapse_vectorized_loop ) {
136+ Optional<Array<PrimExpr>> src_base_indices =
137+ ExtractActiveVectorizedLoopBaseIndices (load->indices );
138+ Optional<Array<PrimExpr>> dst_base_indices =
139+ ExtractActiveVectorizedLoopBaseIndices (store->indices );
140+ if (!src_base_indices.defined () || !dst_base_indices.defined ()) {
141+ return Optional<Stmt>();
142+ }
143+ return MakeCPAsyncStmtFromLoads (
144+ store, ptr_info.value (),
145+ /* dst_base_load=*/
146+ BufferLoad (store->buffer , dst_base_indices.value ()),
147+ /* src_base_load=*/
148+ BufferLoad (load->buffer , src_base_indices.value ()),
149+ /* bytes=*/ index_info->total_transfer_bytes , predicated,
150+ predicate_value);
151+ }
126152 return MakeCPAsyncStmtFromLoads (
127153 store, ptr_info.value (),
128154 /* dst_base_load=*/ BufferLoad (store->buffer , store->indices ),
129155 /* src_base_load=*/ BufferLoad (load->buffer , load->indices ),
130- /* bytes=*/ index_info->transfer_bytes , predicated, predicate_value);
156+ /* bytes=*/ index_info->per_access_transfer_bytes , predicated,
157+ predicate_value);
131158 }
132159
133160 Optional<Array<PrimExpr>> src_base_indices =
@@ -147,7 +174,8 @@ class PTXAsyncCopyInjector : public StmtMutator {
147174 store, ptr_info.value (),
148175 /* dst_base_load=*/ BufferLoad (store->buffer , dst_base_indices.value ()),
149176 /* src_base_load=*/ BufferLoad (load->buffer , src_base_indices.value ()),
150- /* bytes=*/ index_info->transfer_bytes , predicated, predicate_value);
177+ /* bytes=*/ index_info->per_access_transfer_bytes , predicated,
178+ predicate_value);
151179 }
152180
153181 Stmt VisitStmt_ (const SeqStmtNode *op) final {
@@ -301,7 +329,9 @@ class PTXAsyncCopyInjector : public StmtMutator {
301329 PrimExpr src_index;
302330 PrimExpr dst_index;
303331 int index_lanes{1 };
304- int transfer_bytes{0 };
332+ int per_access_transfer_bytes{0 };
333+ int total_transfer_bytes{0 };
334+ bool collapse_vectorized_loop{false };
305335 };
306336
307337 // Pointer element type metadata extracted from buffer handle annotations.
@@ -409,9 +439,17 @@ class PTXAsyncCopyInjector : public StmtMutator {
409439 }
410440
411441 const int effective_lanes = std::max (value_lanes, index_lanes);
412- const int elem_bytes = effective_lanes * load->dtype .bytes ();
413- const int total_bytes = static_cast <int >(elem_bytes) *
414- static_cast <int >(current_vectorized_lanes_);
442+ const int elem_bits = effective_lanes * load->dtype .bits ();
443+ const int total_bits = static_cast <int >(elem_bits) *
444+ static_cast <int >(current_vectorized_lanes_);
445+ // cp.async is byte-granular. We only fold an active vectorized copy into a
446+ // single packed async transfer when the logical payload is exactly
447+ // byte-aligned; otherwise rounding up here would over-copy packed subbyte
448+ // data and change the copy semantics.
449+ if (total_bits % 8 != 0 ) {
450+ return std::nullopt ;
451+ }
452+ const int total_bytes = total_bits / 8 ;
415453 if (!IsValidCPAsyncTransferBytes (total_bytes)) {
416454 return std::nullopt ;
417455 }
@@ -420,10 +458,27 @@ class PTXAsyncCopyInjector : public StmtMutator {
420458 info.src_index = src_index;
421459 info.dst_index = dst_index;
422460 info.index_lanes = index_lanes;
423- info.transfer_bytes = elem_bytes;
461+ info.per_access_transfer_bytes = (elem_bits + 7 ) / 8 ;
462+ info.total_transfer_bytes = total_bytes;
463+ info.collapse_vectorized_loop =
464+ current_vectorized_lanes_ > 1 && index_lanes == 1 ;
424465 return info;
425466 }
426467
468+ Optional<Array<PrimExpr>>
469+ ExtractActiveVectorizedLoopBaseIndices (const Array<PrimExpr> &indices) {
470+ Array<PrimExpr> base_indices;
471+ base_indices.reserve (indices.size ());
472+ for (PrimExpr index : indices) {
473+ for (const auto &loop : active_vectorized_loops_) {
474+ index = analyzer_.Simplify (Substitute (
475+ index, {{loop.loop_var , IntImm (loop.loop_var ->dtype , 0 )}}));
476+ }
477+ base_indices.push_back (index);
478+ }
479+ return Optional<Array<PrimExpr>>(base_indices);
480+ }
481+
427482 static std::optional<PointerTypeInfo>
428483 PreparePointerTypeInfo (const BufferLoadNode *load,
429484 const BufferStoreNode *store) {
@@ -535,6 +590,28 @@ class PTXAsyncCopyInjector : public StmtMutator {
535590 {IntImm (DataType::Int (32 ), n)}));
536591 }
537592
593+ static bool IsCPAsyncStmt (const Stmt &stmt) {
594+ const auto *eval = stmt.as <EvaluateNode>();
595+ if (eval == nullptr ) {
596+ return false ;
597+ }
598+ const auto *call = eval->value .as <CallNode>();
599+ if (call == nullptr ) {
600+ return false ;
601+ }
602+ return call->op .same_as (builtin::ptx_cp_async ()) ||
603+ call->op .same_as (tl::ptx_cp_async ());
604+ }
605+
606+ static bool CanCollapseVectorizedCPAsyncLoop (const Stmt &stmt,
607+ const Var &loop_var) {
608+ if (!IsCPAsyncStmt (stmt)) {
609+ return false ;
610+ }
611+ return !tir::UsesVar (
612+ stmt, [loop_var](const VarNode *v) { return v == loop_var.get (); });
613+ }
614+
538615 // ---- Vectorized-offset contiguity helpers ----
539616 static bool TryGetConstInt64 (const PrimExpr &expr, int64_t *value) {
540617 if (const auto *imm = expr.as <IntImmNode>()) {
0 commit comments