From 428bfb7e1c931715cd3d74f9c65e9991afd86df8 Mon Sep 17 00:00:00 2001 From: Oleh Prypin Date: Mon, 5 Oct 2026 15:07:23 -0700 Subject: [PATCH] Replace `# pytype: disable` suppressions with `# pyrefly: ignore` PiperOrigin-RevId: 993927721 --- examples/policy_improvement_demo.py | 2 -- examples/visualization_demo.py | 12 ++++++------ mctx/_src/policies.py | 2 +- mctx/_src/search.py | 10 ++++------ mctx/_src/tests/tree_test.py | 2 +- mctx/_src/tree.py | 6 ++---- 6 files changed, 14 insertions(+), 20 deletions(-) diff --git a/examples/policy_improvement_demo.py b/examples/policy_improvement_demo.py index 4784749..9e66d00 100644 --- a/examples/policy_improvement_demo.py +++ b/examples/policy_improvement_demo.py @@ -136,10 +136,8 @@ def main(_): # Printing the obtained increase of the policy value. # The obtained increase should be non-negative. action_value_improvement = ( - # pyrefly: ignore[unsupported-operation] output.selected_action_value - output.prior_policy_action_value) weights_value_improvement = ( - # pyrefly: ignore[unsupported-operation] output.action_weights_policy_value - output.prior_policy_value) print("action value improvement: %.3f (min=%.3f)" % (action_value_improvement.mean(), action_value_improvement.min())) diff --git a/examples/visualization_demo.py b/examples/visualization_demo.py index 0af450b..dd9eb84 100644 --- a/examples/visualization_demo.py +++ b/examples/visualization_demo.py @@ -72,7 +72,7 @@ def edge_to_str(node_i, a_i): probs = jax.nn.softmax(tree.children_prior_logits[batch_index, node_i]) # pyrefly: ignore[bad-index] # pyrefly: ignore[unsupported-operation] return (f"{action_labels[a_i]}\n" - f"Q: {tree.qvalues(node_index)[batch_index, a_i]:.2f}\n" # pytype: disable=unsupported-operands # always-use-return-annotations + f"Q: {tree.qvalues(node_index)[batch_index, a_i]:.2f}\n" f"p: {probs[a_i]:.2f}\n") graph = pygraphviz.AGraph(directed=True) @@ -179,11 +179,11 @@ def recurrent_fn(params, rng_key, action, embedding): chex.assert_shape(action, [batch_size]) chex.assert_shape(embedding, [batch_size]) recurrent_fn_output = mctx.RecurrentFnOutput( - reward=rewards[embedding, action], # pyrefly: ignore[bad-index] - discount=discounts[embedding, action], # pyrefly: ignore[bad-index] - prior_logits=prior_logits[embedding], # pyrefly: ignore[bad-index] - value=values[embedding]) # pyrefly: ignore[bad-index] - next_embedding = transition_matrix[embedding, action] # pyrefly: ignore[bad-index] + reward=rewards[embedding, action], + discount=discounts[embedding, action], + prior_logits=prior_logits[embedding], + value=values[embedding]) + next_embedding = transition_matrix[embedding, action] return recurrent_fn_output, next_embedding return root, recurrent_fn diff --git a/mctx/_src/policies.py b/mctx/_src/policies.py index 2266305..908eb86 100644 --- a/mctx/_src/policies.py +++ b/mctx/_src/policies.py @@ -214,7 +214,7 @@ def gumbel_muzero_policy( # a smaller number of valid actions. considered_visit = jnp.max(summary.visit_counts, axis=-1, keepdims=True) # The completed_qvalues include imputed values for unvisited actions. - completed_qvalues = jax.vmap(qtransform, in_axes=[0, None])( # pytype: disable=wrong-arg-types # numpy-scalars # pylint: disable=line-too-long + completed_qvalues = jax.vmap(qtransform, in_axes=[0, None])( # numpy-scalars # pylint: disable=line-too-long search_tree, search_tree.ROOT_INDEX) # pyrefly: ignore[bad-argument-type] to_argmax = seq_halving.score_considered( considered_visit, gumbel, root.prior_logits, completed_qvalues, diff --git a/mctx/_src/search.py b/mctx/_src/search.py index 389d266..e844c73 100644 --- a/mctx/_src/search.py +++ b/mctx/_src/search.py @@ -163,7 +163,7 @@ def body_fun(state): is_before_depth_cutoff = depth < max_depth is_visited = next_node_index != Tree.UNVISITED is_continuing = jnp.logical_and(is_visited, is_before_depth_cutoff) - return _SimulationState( # pytype: disable=wrong-arg-types # jax-types + return _SimulationState( rng_key=rng_key, node_index=node_index, action=action, # pyrefly: ignore[bad-argument-type] @@ -173,15 +173,13 @@ def body_fun(state): node_index = jnp.array(Tree.ROOT_INDEX, dtype=jnp.int32) depth = jnp.zeros((), dtype=tree.children_prior_logits.dtype) - # pytype: disable=wrong-arg-types # jnp-type initial_state = _SimulationState( rng_key=rng_key, node_index=tree.NO_PARENT, action=tree.NO_PARENT, - next_node_index=node_index, - depth=depth, - is_continuing=jnp.array(True)) - # pytype: enable=wrong-arg-types + next_node_index=node_index, # pyrefly: ignore[bad-argument-type] + depth=depth, # pyrefly: ignore[bad-argument-type] + is_continuing=jnp.array(True)) # pyrefly: ignore[bad-argument-type] end_state = jax.lax.while_loop(cond_fun, body_fun, initial_state) # Returning a node with a selected action. diff --git a/mctx/_src/tests/tree_test.py b/mctx/_src/tests/tree_test.py index 6b07d4a..31c7986 100644 --- a/mctx/_src/tests/tree_test.py +++ b/mctx/_src/tests/tree_test.py @@ -115,7 +115,7 @@ def tree_to_pytree(tree: mctx.Tree, batch_i: int = 0): else: child = _create_bare_pynode(prior=prior, action=a_i) # pylint: disable=line-too-long - nodes[node_i]["child_stats"].append(child) # pytype: disable=attribute-error + nodes[node_i]["child_stats"].append(child) # pylint: enable=line-too-long return nodes[0] diff --git a/mctx/_src/tree.py b/mctx/_src/tree.py index e159c82..c2baea5 100644 --- a/mctx/_src/tree.py +++ b/mctx/_src/tree.py @@ -87,12 +87,10 @@ def num_simulations(self): def qvalues(self, indices): """Compute q-values for any node indices in the tree.""" - # pytype: disable=wrong-arg-types # jnp-type if jnp.asarray(indices).shape: return jax.vmap(_unbatched_qvalues)(self, indices) else: return _unbatched_qvalues(self, indices) - # pytype: enable=wrong-arg-types def summary(self) -> SearchSummary: """Extract summary statistics for the root node.""" @@ -109,7 +107,7 @@ def summary(self) -> SearchSummary: visit_probs = visit_counts / jnp.maximum(total_counts, 1) visit_probs = jnp.where(total_counts > 0, visit_probs, 1 / self.num_actions) # Return relevant stats. - return SearchSummary( # pytype: disable=wrong-arg-types # numpy-scalars + return SearchSummary( visit_counts=visit_counts, visit_probs=visit_probs, value=value, @@ -137,7 +135,7 @@ class SearchSummary: def _unbatched_qvalues(tree: Tree, index: int) -> int: chex.assert_rank(tree.children_discounts, 2) - return ( # pytype: disable=bad-return-type # numpy-scalars + return ( # pyrefly: ignore[bad-index, bad-return] tree.children_rewards[index] # pyrefly: ignore[bad-index]