jax.tree 模块#
用于处理树状容器数据结构的实用工具。
jax.tree 命名空间包含了来自 jax.tree_util 的实用工具别名。
函数列表#
|
对树的叶子节点调用 all()。 |
|
将树前缀广播到给定树的完整结构中。 |
|
展平一个 pytree。 |
|
像 |
|
获取 pytree 的叶子节点。 |
|
像 |
|
在 pytree 参数上映射一个多输入函数以生成一个新的 pytree。 |
|
在 pytree 键路径和参数上映射一个多输入函数以生成一个新的 pytree。 |
|
对树的叶子节点调用 reduce()。 |
|
使用关联二元运算对 pytree 执行归约操作。 |
|
声明静态 pytree 属性的便利包装器。 |
|
获取 pytree 的 treedef。 |
|
将具有 (外部, 内部) 结构的树转换为具有 (内部, 外部) 结构的树。 |
|
根据 treedef 和叶子节点重构 pytree。 |