Skip to content
Open
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
55 changes: 37 additions & 18 deletions neuroscope/tracers/jax.py
Original file line number Diff line number Diff line change
Expand Up @@ -214,22 +214,53 @@ def _parse_jaxpr(self, closed_jaxpr: Any) -> None:
def _process_equation(self, eqn: Any, var_to_node: dict[str, str]) -> None:
"""Process a single jaxpr equation into a node."""
primitive_name = eqn.primitive.name
handler_name = f"_handle_{primitive_name}"
handler = getattr(self, handler_name, self._default_handle_equation)
handler(eqn, var_to_node)

def _default_handle_equation(self, eqn: Any, var_to_node: dict[str, str]) -> None:
"""Default handler for jaxpr equations."""
primitive_name = eqn.primitive.name
node_id = f"op_{self._execution_order}"

# Get input tensors
input_tensors = self._get_input_tensors(eqn)
output_tensors = self._get_output_tensors(eqn, node_id, var_to_node)
extra_info = self._extract_safe_params(eqn)

node = GraphNode(
id=node_id,
label=primitive_name,
module_type=f"jax.lax.{primitive_name}",
node_type=NodeType(self._classify_jax_primitive(primitive_name)),
input_tensors=input_tensors,
output_tensors=output_tensors,
execution_order=self._execution_order,
extra_info=extra_info,
)
self._graph.add_node(node)

self._create_edges(eqn, node_id, var_to_node)
self._execution_order += 1

def _get_input_tensors(self, eqn: Any) -> list[TensorMetadata]:
"""Extract input tensors from a jaxpr equation."""
input_tensors = []
for invar in eqn.invars:
if hasattr(invar, "aval"):
input_tensors.append(self._aval_to_metadata(invar.aval))
return input_tensors

# Get output tensors
def _get_output_tensors(self, eqn: Any, node_id: str, var_to_node: dict[str, str]) -> list[TensorMetadata]:
"""Extract output tensors from a jaxpr equation and update var_to_node mapping."""
output_tensors = []
for outvar in eqn.outvars:
if hasattr(outvar, "aval"):
output_tensors.append(self._aval_to_metadata(outvar.aval))
var_to_node[str(outvar)] = node_id
return output_tensors

# Extract safe params (skip non-serializable ones)
def _extract_safe_params(self, eqn: Any) -> dict[str, Any]:
"""Extract serializable parameters from a jaxpr equation."""
extra_info = {}
if eqn.params:
for k, v in eqn.params.items():
Expand All @@ -241,20 +272,10 @@ def _process_equation(self, eqn: Any, var_to_node: dict[str, str]) -> None:
extra_info[k] = str(v)
except Exception:
pass
return extra_info

node = GraphNode(
id=node_id,
label=primitive_name,
module_type=f"jax.lax.{primitive_name}",
node_type=NodeType(self._classify_jax_primitive(primitive_name)),
input_tensors=input_tensors,
output_tensors=output_tensors,
execution_order=self._execution_order,
extra_info=extra_info,
)
self._graph.add_node(node)

# Create edges from inputs
def _create_edges(self, eqn: Any, node_id: str, var_to_node: dict[str, str]) -> None:
"""Create edges from inputs to the current node."""
for i, invar in enumerate(eqn.invars):
source_id = var_to_node.get(str(invar))
if source_id and source_id != node_id:
Expand All @@ -265,8 +286,6 @@ def _process_equation(self, eqn: Any, var_to_node: dict[str, str]) -> None:
)
self._graph.add_edge(edge)

self._execution_order += 1

def _aval_to_metadata(self, aval: Any) -> TensorMetadata:
"""Convert JAX abstract value to TensorMetadata."""
try:
Expand Down