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
2 changes: 0 additions & 2 deletions examples/policy_improvement_demo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()))
Expand Down
12 changes: 6 additions & 6 deletions examples/visualization_demo.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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
Expand Down
2 changes: 1 addition & 1 deletion mctx/_src/policies.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
10 changes: 4 additions & 6 deletions mctx/_src/search.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]
Expand All @@ -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.
Expand Down
2 changes: 1 addition & 1 deletion mctx/_src/tests/tree_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]

Expand Down
6 changes: 2 additions & 4 deletions mctx/_src/tree.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""
Expand All @@ -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,
Expand Down Expand Up @@ -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]
Expand Down
Loading