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


##########
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:
   Good catch, yes we should



-- 
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