jax
9152b760 - Add with_sharding_constraint method to be used within sharded_jit. (#3100)

Commit
6 years ago
Add with_sharding_constraint method to be used within sharded_jit. (#3100) See the with_sharding_constraint docstring for a description of what this method does. Depending on how we decide nested sharded_jits should work, an alternative implementation for with_sharding_constraint could be: ```python def with_sharding_constraint(x, partitions): return sharded_jit(lambda x: x, in_parts=partitions, out_parts=partitions) ``` In this case, we could get rid of the with_sharding_constraint primitive, and possibly even the API. This implementation gets the job done for now without committing to a nested sharded_jit behavior, and is also much easier to take the gradient of than sharded_jit.
Author
Parents
Loading