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.