Decorator factory to enable compiling of a function if the minimum PyTorch version requirement is met. Args: min_version (str, optional): Minimum PyTorch version required (e.g., "2.7.0"). If None, the function is always enabled. Returns: Callable: A decorat
(min_version=None)
| 40 | |
| 41 | |
| 42 | def enable(min_version=None): |
| 43 | """ |
| 44 | Decorator factory to enable compiling of a function if the minimum PyTorch version requirement is met. |
| 45 | |
| 46 | Args: |
| 47 | min_version (str, optional): Minimum PyTorch version required (e.g., "2.7.0"). |
| 48 | If None, the function is always enabled. |
| 49 | |
| 50 | Returns: |
| 51 | Callable: A decorator that wraps the function. |
| 52 | |
| 53 | Examples: |
| 54 | @enable("2.7.0") |
| 55 | def my_function(): |
| 56 | pass |
| 57 | |
| 58 | @enable |
| 59 | def another_function(): |
| 60 | pass |
| 61 | """ |
| 62 | |
| 63 | def decorator(func): |
| 64 | if not is_compiling(): |
| 65 | return func |
| 66 | |
| 67 | @functools.wraps(func) |
| 68 | def wrapper(*args, **kwargs): |
| 69 | if min_version is None or required_torch_version(min_version=min_version): |
| 70 | return func(*args, **kwargs) |
| 71 | return disable(func)(*args, **kwargs) |
| 72 | |
| 73 | return wrapper |
| 74 | |
| 75 | # Called with no arguments |
| 76 | if callable(min_version): |
| 77 | func = min_version |
| 78 | min_version = None |
| 79 | return decorator(func) |
| 80 | |
| 81 | return decorator |
| 82 | |
| 83 | |
| 84 | def is_compiling(): |