Files
pytorch/functorch/dim/batch_tensor.py
Xuehai Pan 740fb22966 [BE][Easy][4/19] enforce style for empty lines in import segments in functorch/ (#129755)
See https://github.com/pytorch/pytorch/pull/129751#issue-2380881501. Most changes are auto-generated by linter.

You can review these PRs via:

```bash
git diff --ignore-all-space --ignore-blank-lines HEAD~1
```

Pull Request resolved: https://github.com/pytorch/pytorch/pull/129755
Approved by: https://github.com/zou3519
ghstack dependencies: #129752
2024-07-18 05:08:03 +00:00

27 lines
668 B
Python

# Copyright (c) Facebook, Inc. and its affiliates.
# All rights reserved.
#
# This source code is licensed under the BSD-style license found in the
# LICENSE file in the root directory of this source tree.
from contextlib import contextmanager
from torch._C._functorch import _vmap_add_layers, _vmap_remove_layers
_enabled = False
@contextmanager
def _enable_layers(dims):
global _enabled
assert not _enabled
input = sorted((d._level, d.size) for d in dims if not isinstance(d, int))
n = len(input)
try:
_vmap_add_layers(input)
_enabled = True
yield
finally:
_enabled = False
_vmap_remove_layers(n)