(self)
| 762 | "change": expected_subset, |
| 763 | }, |
| 764 | ) |
| 765 | |
| 766 | def test_complete_partition_spec_tree(self): |
| 767 | data = dict( |
| 768 | replicated=dict(a=1, b=2), |
| 769 | sharded=VDict(c=3, d=4), |
| 770 | ) |
| 771 | partition_by_x = PartitionSpec("x") |
| 772 | partial_partition_spec = dict(replicated=None, sharded=partition_by_x) |
| 773 | self.assertEqual( |
| 774 | complete_partition_spec_tree( |
| 775 | jax.tree_util.tree_structure(data), partial_partition_spec |
| 776 | ), |
| 777 | dict( |
| 778 | replicated=dict(a=None, b=None), sharded=VDict(c=partition_by_x, d=partition_by_x) |
| 779 | ), |
| 780 | ) |
| 781 | param_spec = ParameterSpec( |
| 782 | shape=[1, 2, 3], |
| 783 | mesh_axes=["x", "y", "z"], |
| 784 | factorization=FactorizationSpec(axes=[None, "row", "col"]), |
| 785 | ) |
| 786 | self.assertEqual( |
| 787 | complete_partition_spec_tree( |
| 788 | jax.tree_util.tree_structure(data), dict(replicated=None, sharded=param_spec) |
| 789 | ), |
| 790 | dict(replicated=dict(a=None, b=None), sharded=VDict(c=param_spec, d=param_spec)), |
| 791 | ) |
| 792 |
nothing calls this directly
no test coverage detected