jax.tree 模块

目录

jax.tree 模块#

用于处理树状容器数据结构的实用工具。

jax.tree 命名空间包含了来自 jax.tree_util 的实用工具别名。

函数列表#

all(tree, *[, is_leaf])

对树的叶子节点调用 all()。

broadcast(prefix_tree, full_tree[, is_leaf])

将树前缀广播到给定树的完整结构中。

flatten(tree[, is_leaf])

展平一个 pytree。

flatten_with_path(tree[, is_leaf, ...])

tree_flatten 一样展平 pytree,但同时返回每个叶子节点的键路径。

leaves(tree[, is_leaf])

获取 pytree 的叶子节点。

leaves_with_path(tree[, is_leaf, ...])

tree_leaves 一样获取 pytree 的叶子节点,并返回每个叶子节点的键路径。

map(f, tree, *rest[, is_leaf])

在 pytree 参数上映射一个多输入函数以生成一个新的 pytree。

map_with_path(f, tree, *rest[, is_leaf, ...])

在 pytree 键路径和参数上映射一个多输入函数以生成一个新的 pytree。

reduce(function, tree[, initializer, is_leaf])

对树的叶子节点调用 reduce()。

reduce_associative(operation, tree, *[, ...])

使用关联二元运算对 pytree 执行归约操作。

static(**kwargs)

声明静态 pytree 属性的便利包装器。

structure(tree[, is_leaf])

获取 pytree 的 treedef。

transpose(outer_treedef, inner_treedef, ...)

将具有 (外部, 内部) 结构的树转换为具有 (内部, 外部) 结构的树。

unflatten(treedef, leaves)

根据 treedef 和叶子节点重构 pytree。