From 92f92a98e445acbfddfd1e294e994fa9c7f08a1b Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Thu, 25 Sep 2025 21:40:07 -0400 Subject: [PATCH 1/2] [DEV] Add support for join operation in sanitizer - Add Join class to data.py for representing join operations - Import Join in patch.py and sanitizer.py - Add mappings for create_join in patch.py - Implement op_join_overrider in sanitizer to handle join operations - Add join to BROADCAST_OPS and create spec with proper dtype handling - Register Join in OP_TYPE_TO_OVERRIDER mapping - Add NotImplementedError for join in Z3 evaluation (to be implemented later) --- triton_viz/clients/sanitizer/sanitizer.py | 23 ++++++++++++++++++++++- triton_viz/core/data.py | 5 +++++ triton_viz/core/patch.py | 4 ++++ 3 files changed, 31 insertions(+), 1 deletion(-) diff --git a/triton_viz/clients/sanitizer/sanitizer.py b/triton_viz/clients/sanitizer/sanitizer.py index 06f475742..f24799153 100644 --- a/triton_viz/clients/sanitizer/sanitizer.py +++ b/triton_viz/clients/sanitizer/sanitizer.py @@ -55,6 +55,7 @@ Rsqrt, CastImpl, Reshape, + Join, Fabs, Ashr, Advance, @@ -628,6 +629,11 @@ def _broadcast_dtype(self): self.dtype_tt = self.children["arg"].dtype_tt +def _join_post(self): + # Join preserves the dtype of the left hand side + self.dtype_tt = self.children["lhs"].dtype_tt + + def _binary_dtype(expr): expr.dtype_tt = expr.lhs.dtype_tt @@ -706,7 +712,7 @@ class SymbolicExpr: REDUCE_OPS = ("sum", "max", "min", "dot") SCAN_OPS = ("cumsum",) POINTER_OPS = ("make_block_ptr", "addptr", "advance") - BROADCAST_OPS = ("splat", "expand_dims", "broadcast", "reshape") + BROADCAST_OPS = ("splat", "expand_dims", "broadcast", "reshape", "join") CAST_OPS = ("cast_impl", "bitcast") ATOMIC_OPS = ("atomic_cas",) SUPPORTED_OPS = ( @@ -791,6 +797,7 @@ class SymbolicExpr: "expand_dims": Spec(req=("arg", "axis"), post=_broadcast_dtype), "broadcast": Spec(req=("arg", "shape"), post=_broadcast_dtype), "reshape": Spec(req=("arg", "shape"), post=_broadcast_dtype), + "join": Spec(req=("lhs", "rhs"), post=_join_post), # Casting "cast_impl": Spec(req=("src", "dst_type"), post=_cast_impl_post), "bitcast": Spec(req=("src", "dst_type"), post=_cast_impl_post), @@ -1178,6 +1185,9 @@ def _to_z3(self) -> tuple[ArithRef, list]: if self.op in ("splat", "expand_dims", "broadcast"): self._z3, self._constraints = self.arg._to_z3() + if self.op == "join": + raise NotImplementedError("Join operation is not implemented in Z3 evaluation yet") + if self.op == "addptr": # Add pointer operation ptr_z3, constraints_ptr = self.ptr._to_z3() @@ -1890,6 +1900,16 @@ def op_reshape_overrider(arg, shape, allow_reorder): result.dtype_tt = arg_sym.dtype_tt if hasattr(arg_sym, "dtype_tt") else None return result + def op_join_overrider(lhs, rhs): + # Join operation combines two tensors along the last axis + lhs_sym = SymbolicExpr.from_value(lhs) + rhs_sym = SymbolicExpr.from_value(rhs) + # Create a join symbolic expression + result = SymbolicExpr("join", lhs_sym, rhs_sym) + # The dtype should be preserved from the inputs + result.dtype_tt = lhs_sym.dtype_tt if hasattr(lhs_sym, "dtype_tt") else None + return result + def op_fabs_overrider(arg): arg_sym = SymbolicExpr.from_value(arg) return SymbolicExpr("fabs", arg_sym) @@ -1963,6 +1983,7 @@ def op_atomic_cas_overrider(ptr, cmp, val, sem, scope): Rsqrt: op_rsqrt_overrider, CastImpl: op_cast_impl_overrider, Reshape: op_reshape_overrider, + Join: op_join_overrider, Fabs: op_fabs_overrider, Ashr: op_ashr_overrider, Advance: op_advance_overrider, diff --git a/triton_viz/core/data.py b/triton_viz/core/data.py index 1c2e9b55f..411820488 100644 --- a/triton_viz/core/data.py +++ b/triton_viz/core/data.py @@ -188,6 +188,11 @@ class Reshape(Op): name: ClassVar[str] = "reshape" +@dataclass +class Join(Op): + name: ClassVar[str] = "join" + + @dataclass class Fabs(Op): name: ClassVar[str] = "fabs" diff --git a/triton_viz/core/patch.py b/triton_viz/core/patch.py index 09ba5aa67..af7288a64 100644 --- a/triton_viz/core/patch.py +++ b/triton_viz/core/patch.py @@ -32,6 +32,7 @@ Rsqrt, CastImpl, Reshape, + Join, Fabs, Ashr, Advance, @@ -78,6 +79,7 @@ Rsqrt, CastImpl, Reshape, + Join, Fabs, Ashr, Advance, @@ -112,6 +114,7 @@ Rsqrt: "create_rsqrt", CastImpl: "cast_impl", Reshape: "create_reshape", + Join: "create_join", Fabs: "create_fabs", Ashr: "create_ashr", Advance: "create_advance", @@ -144,6 +147,7 @@ Rsqrt: interpreter_builder.create_rsqrt, CastImpl: interpreter_builder.cast_impl, Reshape: interpreter_builder.create_reshape, + Join: interpreter_builder.create_join, Fabs: interpreter_builder.create_fabs, Ashr: interpreter_builder.create_ashr, Advance: interpreter_builder.create_advance, From eb7b42c2f55324829beed367fac00df41c4360cf Mon Sep 17 00:00:00 2001 From: Hao Wu Date: Thu, 25 Sep 2025 21:56:29 -0400 Subject: [PATCH 2/2] pre-commit formatting --- triton_viz/clients/sanitizer/sanitizer.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/triton_viz/clients/sanitizer/sanitizer.py b/triton_viz/clients/sanitizer/sanitizer.py index f24799153..5bda5f3df 100644 --- a/triton_viz/clients/sanitizer/sanitizer.py +++ b/triton_viz/clients/sanitizer/sanitizer.py @@ -1186,7 +1186,9 @@ def _to_z3(self) -> tuple[ArithRef, list]: self._z3, self._constraints = self.arg._to_z3() if self.op == "join": - raise NotImplementedError("Join operation is not implemented in Z3 evaluation yet") + raise NotImplementedError( + "Join operation is not implemented in Z3 evaluation yet" + ) if self.op == "addptr": # Add pointer operation