Tests sharding and unsharding scalars.
(self)
| 127 | _ = p.get_unsharded_shape([[2], [4, 3]]) |
| 128 | |
| 129 | def testScalar(self): |
| 130 | """Tests sharding and unsharding scalars.""" |
| 131 | p = tpu_sharding.ShardingPolicy() |
| 132 | p.freeze() |
| 133 | self.assertEqual(p.get_sharded_shape([]), []) |
| 134 | self.assertEqual(p.get_unsharded_shape([[]]), []) |
| 135 | |
| 136 | |
| 137 | if __name__ == "__main__": |
nothing calls this directly
no test coverage detected