Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 24 additions & 1 deletion triton_viz/clients/sanitizer/sanitizer.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@
Rsqrt,
CastImpl,
Reshape,
Join,
Fabs,
Ashr,
Advance,
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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 = (
Expand Down Expand Up @@ -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),
Expand Down Expand Up @@ -1178,6 +1185,11 @@ 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()
Expand Down Expand Up @@ -1890,6 +1902,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)
Expand Down Expand Up @@ -1963,6 +1985,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,
Expand Down
5 changes: 5 additions & 0 deletions triton_viz/core/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
4 changes: 4 additions & 0 deletions triton_viz/core/patch.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
Rsqrt,
CastImpl,
Reshape,
Join,
Fabs,
Ashr,
Advance,
Expand Down Expand Up @@ -78,6 +79,7 @@
Rsqrt,
CastImpl,
Reshape,
Join,
Fabs,
Ashr,
Advance,
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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,
Expand Down