ekalda commented on code in PR #16966:
URL: https://github.com/apache/tvm/pull/16966#discussion_r1605010649


##########
include/tvm/tir/buffer.h:
##########
@@ -209,14 +209,20 @@ class Buffer : public ObjectRef {
    * \brief Create an Expr that does a vector load at begin index.
    * \param begin The beginning index
    * \param dtype The data type to be loaded.
+   * \param predicate A vector mask of boolean values indicating which lanes 
of a vector are to be
+   * stored. The number lanes of the mask must be equal to the number of lanes 
in value.

Review Comment:
   Nit: "storing lanes of a vector" is not relevant to a loading API 
documentation, also there is no `value` argument here, so it's a bit ambiguous. 
(Same for other documentation of loading in other places)



##########
src/tir/ir/stmt.cc:
##########
@@ -476,29 +477,39 @@ BufferStore::BufferStore(Buffer buffer, PrimExpr value, 
Array<PrimExpr> indices,
   ICHECK(!(is_index_scalable && is_buffer_dtype_scalable))
       << "Index dtype and buffer dtype can't both be scalable.";
 
-  if (is_index_scalable || is_buffer_dtype_scalable) {
-    ICHECK(is_value_dtype_scalable) << "Can't store non-scalable data into 
scalable buffer";
+  if (predicate.defined()) {
+    bool is_predicate_dtype_scalable = 
predicate.value().dtype().is_scalable_vector();
+    ICHECK_EQ(is_value_dtype_scalable, is_predicate_dtype_scalable)
+        << "Predicate mask dtype and value dtype must both be scalable.";
   }
 
-  int index_lanes;
-  if (indices.empty()) {
-    index_lanes = 1;
-  } else if (is_index_scalable) {
-    index_lanes = indices.back().dtype().vscale_factor();
-  } else {
-    index_lanes = indices.back().dtype().lanes();
+  if (is_index_scalable || is_buffer_dtype_scalable) {
+    ICHECK(is_value_dtype_scalable) << "Can't store non-scalable data into 
scalable buffer";
   }
 
-  int buffer_lanes =
-      is_buffer_dtype_scalable ? buffer->dtype.vscale_factor() : 
buffer->dtype.lanes();
-  int value_dtype_lanes =
-      is_value_dtype_scalable ? value.dtype().vscale_factor() : 
value.dtype().lanes();
+  int index_lanes = indices.empty() ? 1 : 
indices.back().dtype().get_lanes_or_vscale_factor();
+  int buffer_lanes = buffer->dtype.get_lanes_or_vscale_factor();
+  int value_dtype_lanes = value.dtype().get_lanes_or_vscale_factor();
 
   ICHECK_EQ(index_lanes * buffer_lanes, value_dtype_lanes)
       << "Cannot store value with " << value_dtype_lanes << ", expected value 
with "
       << index_lanes * buffer_lanes << " (" << index_lanes << " index lanes * 
" << buffer_lanes
       << " buffer element lanes)";
 
+  if (predicate.defined()) {
+    DataType predicate_dtype = predicate.value().dtype();
+    int predicate_dtype_lanes = predicate_dtype.get_lanes_or_vscale_factor();
+    ICHECK_EQ(value_dtype_lanes, predicate_dtype_lanes)

Review Comment:
   Should we have a similar check in the `BufferLoad` node? 



##########
src/tir/transforms/inject_rolling_buffer.cc:
##########
@@ -257,7 +257,9 @@ class RollingBufferInjector : public StmtExprMutator {
           indices.push_back(index);
         }
       }
-      Stmt buffer_store = BufferStore(op->buffer, op->value, indices, 
op->span);
+      ICHECK(!op->predicate.defined())
+          << "Predicated buffer store is not current supported in the inject 
rolling buffer pass.";

Review Comment:
   Nit:
   
   ```suggestion
             << "Predicated buffer store is not currently supported in the 
inject rolling buffer pass.";
   ```



-- 
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.

To unsubscribe, e-mail: [email protected]

For queries about this service, please contact Infrastructure at:
[email protected]

Reply via email to