From 63e6d3fcd847d4ef03fc5b77e78331b502680e5b Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Thu, 25 Sep 2025 19:39:37 -0400 Subject: [PATCH 1/3] [FIX] Fix sanitizer handling of multi-element TensorHandle - Update _infer_literal_dtype to handle TensorHandle with multiple data elements - Verify dtype consistency across all elements in multi-element cases - Modify from_value to preserve entire array for multi-element TensorHandle - Maintain backward compatibility for single-element cases This fixes crashes when using `tl.interleave` and similar operations that produce TensorHandle objects with multiple internal values. --- triton_viz/clients/sanitizer/sanitizer.py | 32 +++++++++++++++++++---- 1 file changed, 27 insertions(+), 5 deletions(-) diff --git a/triton_viz/clients/sanitizer/sanitizer.py b/triton_viz/clients/sanitizer/sanitizer.py index 6765da86..bd031af5 100644 --- a/triton_viz/clients/sanitizer/sanitizer.py +++ b/triton_viz/clients/sanitizer/sanitizer.py @@ -997,10 +997,25 @@ def _infer_literal_dtype(var): ) return first_dtype if isinstance(var, TensorHandle): - if len(var.data) != 1: - raise ValueError( - f"Unsupported var.data: {var.data} with length more than one!" - ) + # Handle TensorHandle with potentially multiple data elements + if len(var.data) > 1: + # For multi-element cases (e.g., from interleave operations), + # verify dtype consistency across all elements + first_dtype = None + for elem in var.data: + if hasattr(elem, "dtype"): + elem_dtype = elem.dtype + if first_dtype is None: + first_dtype = elem_dtype + elif first_dtype != elem_dtype: + raise ValueError( + f"TensorHandle contains elements with inconsistent dtypes: " + f"{first_dtype} vs {elem_dtype}" + ) + # If all elements have consistent dtype or no dtype attributes, + # trust and return the TensorHandle's overall dtype + return var.dtype + # Single element case - original logic if var.dtype in SymbolicExpr.triton_scala_dtypes: # if an immediate return var.dtype if isinstance(var.dtype, tl.pointer_type): # if a pointer @@ -1029,7 +1044,14 @@ def from_value(cls, var): if isinstance(var, SymbolicExpr.tuple_types): # if a tuple return cls("const", tuple(var), dtype_tt) if isinstance(var, TensorHandle): # if a TensorHandle - return cls("const", var.data.item(), dtype_tt) + # Handle both single and multi-element TensorHandle + if len(var.data) == 1: + # Single element: extract scalar for backward compatibility + return cls("const", var.data.item(), dtype_tt) + else: + # Multi-element: treat like a tuple, keep the entire array + # This occurs in operations like interleave + return cls("const", var.data, dtype_tt) if isinstance( var, SymbolicExpr.builtin_scala_types ): # if a python builtin type From e6e997cfdd56a2de92df044b189788d2b86733c6 Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Thu, 25 Sep 2025 20:31:48 -0400 Subject: [PATCH 2/3] [DEV] Add support for bitwise XOR and OR operations in sanitizer - Add bitwise_xor (^) and bitwise_or (|) to BINARY_OP_SYMBOL_TABLE - Add handlers for np.bitwise_xor and np.bitwise_or in op_binary_op_overrider - Include bitwise_xor and bitwise_or in OP_SPEC binary operations list This fixes crashes when using bitwise XOR/OR operations in Triton kernels, particularly needed for FP4 to BF16 conversion code. --- triton_viz/clients/sanitizer/sanitizer.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/triton_viz/clients/sanitizer/sanitizer.py b/triton_viz/clients/sanitizer/sanitizer.py index 7827007e..2355b935 100644 --- a/triton_viz/clients/sanitizer/sanitizer.py +++ b/triton_viz/clients/sanitizer/sanitizer.py @@ -697,6 +697,8 @@ class SymbolicExpr: "maximum": "max", "minimum": "min", "bitwise_and": "&", + "bitwise_or": "|", + "bitwise_xor": "^", "right_shift": ">>", "left_shift": "<<", "ashr": ">>>", @@ -767,6 +769,8 @@ class SymbolicExpr: "maximum", "minimum", "bitwise_and", + "bitwise_or", + "bitwise_xor", "right_shift", "left_shift", "ashr", @@ -1784,6 +1788,8 @@ def op_binary_op_overrider(lhs, rhs, op): np.maximum: lambda lhs, rhs: SymbolicExpr("maximum", lhs, rhs), np.minimum: lambda lhs, rhs: SymbolicExpr("minimum", lhs, rhs), np.bitwise_and: lambda lhs, rhs: SymbolicExpr("bitwise_and", lhs, rhs), + np.bitwise_or: lambda lhs, rhs: SymbolicExpr("bitwise_or", lhs, rhs), + np.bitwise_xor: lambda lhs, rhs: SymbolicExpr("bitwise_xor", lhs, rhs), np.right_shift: lambda lhs, rhs: SymbolicExpr("right_shift", lhs, rhs), np.left_shift: lambda lhs, rhs: SymbolicExpr("left_shift", lhs, rhs), } From 98bffcf638e7e618d6c58c1b6c5dd8ad3b5d9fea Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Thu, 25 Sep 2025 21:54:59 -0400 Subject: [PATCH 3/3] Revert "[FIX] Fix sanitizer handling of multi-element TensorHandle" This reverts commit 63e6d3fcd847d4ef03fc5b77e78331b502680e5b. --- triton_viz/clients/sanitizer/sanitizer.py | 32 ++++------------------- 1 file changed, 5 insertions(+), 27 deletions(-) diff --git a/triton_viz/clients/sanitizer/sanitizer.py b/triton_viz/clients/sanitizer/sanitizer.py index 2355b935..d6c46305 100644 --- a/triton_viz/clients/sanitizer/sanitizer.py +++ b/triton_viz/clients/sanitizer/sanitizer.py @@ -1006,25 +1006,10 @@ def _infer_literal_dtype(var): ) return first_dtype if isinstance(var, TensorHandle): - # Handle TensorHandle with potentially multiple data elements - if len(var.data) > 1: - # For multi-element cases (e.g., from interleave operations), - # verify dtype consistency across all elements - first_dtype = None - for elem in var.data: - if hasattr(elem, "dtype"): - elem_dtype = elem.dtype - if first_dtype is None: - first_dtype = elem_dtype - elif first_dtype != elem_dtype: - raise ValueError( - f"TensorHandle contains elements with inconsistent dtypes: " - f"{first_dtype} vs {elem_dtype}" - ) - # If all elements have consistent dtype or no dtype attributes, - # trust and return the TensorHandle's overall dtype - return var.dtype - # Single element case - original logic + if len(var.data) != 1: + raise ValueError( + f"Unsupported var.data: {var.data} with length more than one!" + ) if var.dtype in SymbolicExpr.triton_scala_dtypes: # if an immediate return var.dtype if isinstance(var.dtype, tl.pointer_type): # if a pointer @@ -1053,14 +1038,7 @@ def from_value(cls, var): if isinstance(var, SymbolicExpr.tuple_types): # if a tuple return cls("const", tuple(var), dtype_tt) if isinstance(var, TensorHandle): # if a TensorHandle - # Handle both single and multi-element TensorHandle - if len(var.data) == 1: - # Single element: extract scalar for backward compatibility - return cls("const", var.data.item(), dtype_tt) - else: - # Multi-element: treat like a tuple, keep the entire array - # This occurs in operations like interleave - return cls("const", var.data, dtype_tt) + return cls("const", var.data.item(), dtype_tt) if isinstance( var, SymbolicExpr.builtin_scala_types ): # if a python builtin type