diff --git a/neuroscope/tracers/jax.py b/neuroscope/tracers/jax.py index be6da2b..270834a 100644 --- a/neuroscope/tracers/jax.py +++ b/neuroscope/tracers/jax.py @@ -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(): @@ -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: @@ -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: